diff options
| author | Bob Halley <halley@dnspython.org> | 2022-03-15 08:37:20 -0700 |
|---|---|---|
| committer | Bob Halley <halley@dnspython.org> | 2022-03-15 08:37:20 -0700 |
| commit | b1d2332687adbecc0acbb4e623124f783f859d9e (patch) | |
| tree | 5318d5ecc0dd35e0a6922380cd60f9d9caa9ad34 /dns | |
| parent | 08f8bde64e8679d5e4f0b129292461de152ba32b (diff) | |
| download | dnspython-b1d2332687adbecc0acbb4e623124f783f859d9e.tar.gz | |
black autoformatting
Diffstat (limited to 'dns')
106 files changed, 4597 insertions, 2953 deletions
diff --git a/dns/__init__.py b/dns/__init__.py index a620f97..196be22 100644 --- a/dns/__init__.py +++ b/dns/__init__.py @@ -18,51 +18,51 @@ """dnspython DNS toolkit""" __all__ = [ - 'asyncbackend', - 'asyncquery', - 'asyncresolver', - 'dnssec', - 'dnssectypes', - 'e164', - 'edns', - 'entropy', - 'exception', - 'flags', - 'immutable', - 'inet', - 'ipv4', - 'ipv6', - 'message', - 'name', - 'namedict', - 'node', - 'opcode', - 'query', - 'rcode', - 'rdata', - 'rdataclass', - 'rdataset', - 'rdatatype', - 'renderer', - 'resolver', - 'reversename', - 'rrset', - 'serial', - 'set', - 'tokenizer', - 'transaction', - 'tsig', - 'tsigkeyring', - 'ttl', - 'rdtypes', - 'update', - 'version', - 'versioned', - 'wire', - 'xfr', - 'zone', - 'zonetypes', - 'zonefile', + "asyncbackend", + "asyncquery", + "asyncresolver", + "dnssec", + "dnssectypes", + "e164", + "edns", + "entropy", + "exception", + "flags", + "immutable", + "inet", + "ipv4", + "ipv6", + "message", + "name", + "namedict", + "node", + "opcode", + "query", + "rcode", + "rdata", + "rdataclass", + "rdataset", + "rdatatype", + "renderer", + "resolver", + "reversename", + "rrset", + "serial", + "set", + "tokenizer", + "transaction", + "tsig", + "tsigkeyring", + "ttl", + "rdtypes", + "update", + "version", + "versioned", + "wire", + "xfr", + "zone", + "zonetypes", + "zonefile", ] from dns.version import version as __version__ # noqa diff --git a/dns/_asyncbackend.py b/dns/_asyncbackend.py index 674bf6e..ff24604 100644 --- a/dns/_asyncbackend.py +++ b/dns/_asyncbackend.py @@ -3,6 +3,7 @@ # This is a nullcontext for both sync and async. 3.7 has a nullcontext, # but it is only for sync use. + class NullContext: def __init__(self, enter_result=None): self.enter_result = enter_result @@ -23,6 +24,7 @@ class NullContext: # These are declared here so backends can import them without creating # circular dependencies with dns.asyncbackend. + class Socket: # pragma: no cover async def close(self): pass @@ -59,13 +61,21 @@ class StreamSocket(Socket): # pragma: no cover raise NotImplementedError -class Backend: # pragma: no cover +class Backend: # pragma: no cover def name(self): - return 'unknown' - - async def make_socket(self, af, socktype, proto=0, - source=None, destination=None, timeout=None, - ssl_context=None, server_hostname=None): + return "unknown" + + async def make_socket( + self, + af, + socktype, + proto=0, + source=None, + destination=None, + timeout=None, + ssl_context=None, + server_hostname=None, + ): raise NotImplementedError def datagram_connection_required(self): diff --git a/dns/_asyncio_backend.py b/dns/_asyncio_backend.py index 1091777..50bde1d 100644 --- a/dns/_asyncio_backend.py +++ b/dns/_asyncio_backend.py @@ -10,7 +10,8 @@ import dns._asyncbackend import dns.exception -_is_win32 = sys.platform == 'win32' +_is_win32 = sys.platform == "win32" + def _get_running_loop(): try: @@ -76,10 +77,10 @@ class DatagramSocket(dns._asyncbackend.DatagramSocket): self.protocol.close() async def getpeername(self): - return self.transport.get_extra_info('peername') + return self.transport.get_extra_info("peername") async def getsockname(self): - return self.transport.get_extra_info('sockname') + return self.transport.get_extra_info("sockname") class StreamSocket(dns._asyncbackend.StreamSocket): @@ -93,8 +94,7 @@ class StreamSocket(dns._asyncbackend.StreamSocket): return await _maybe_wait_for(self.writer.drain(), timeout) async def recv(self, size, timeout): - return await _maybe_wait_for(self.reader.read(size), - timeout) + return await _maybe_wait_for(self.reader.read(size), timeout) async def close(self): self.writer.close() @@ -104,43 +104,60 @@ class StreamSocket(dns._asyncbackend.StreamSocket): pass async def getpeername(self): - return self.writer.get_extra_info('peername') + return self.writer.get_extra_info("peername") async def getsockname(self): - return self.writer.get_extra_info('sockname') + return self.writer.get_extra_info("sockname") class Backend(dns._asyncbackend.Backend): def name(self): - return 'asyncio' - - async def make_socket(self, af, socktype, proto=0, - source=None, destination=None, timeout=None, - ssl_context=None, server_hostname=None): - if destination is None and socktype == socket.SOCK_DGRAM and \ - _is_win32: - raise NotImplementedError('destinationless datagram sockets ' - 'are not supported by asyncio ' - 'on Windows') + return "asyncio" + + async def make_socket( + self, + af, + socktype, + proto=0, + source=None, + destination=None, + timeout=None, + ssl_context=None, + server_hostname=None, + ): + if destination is None and socktype == socket.SOCK_DGRAM and _is_win32: + raise NotImplementedError( + "destinationless datagram sockets " + "are not supported by asyncio " + "on Windows" + ) loop = _get_running_loop() if socktype == socket.SOCK_DGRAM: transport, protocol = await loop.create_datagram_endpoint( - _DatagramProtocol, source, family=af, - proto=proto, remote_addr=destination) + _DatagramProtocol, + source, + family=af, + proto=proto, + remote_addr=destination, + ) return DatagramSocket(af, transport, protocol) elif socktype == socket.SOCK_STREAM: (r, w) = await _maybe_wait_for( - asyncio.open_connection(destination[0], - destination[1], - ssl=ssl_context, - family=af, - proto=proto, - local_addr=source, - server_hostname=server_hostname), - timeout) + asyncio.open_connection( + destination[0], + destination[1], + ssl=ssl_context, + family=af, + proto=proto, + local_addr=source, + server_hostname=server_hostname, + ), + timeout, + ) return StreamSocket(af, r, w) - raise NotImplementedError('unsupported socket ' + - f'type {socktype}') # pragma: no cover + raise NotImplementedError( + "unsupported socket " + f"type {socktype}" + ) # pragma: no cover async def sleep(self, interval): await asyncio.sleep(interval) diff --git a/dns/_curio_backend.py b/dns/_curio_backend.py index 3f22b5d..765d647 100644 --- a/dns/_curio_backend.py +++ b/dns/_curio_backend.py @@ -32,7 +32,9 @@ class DatagramSocket(dns._asyncbackend.DatagramSocket): async def sendto(self, what, destination, timeout): async with _maybe_timeout(timeout): return await self.socket.sendto(what, destination) - raise dns.exception.Timeout(timeout=timeout) # pragma: no cover lgtm[py/unreachable-statement] + raise dns.exception.Timeout( + timeout=timeout + ) # pragma: no cover lgtm[py/unreachable-statement] async def recvfrom(self, size, timeout): async with _maybe_timeout(timeout): @@ -76,11 +78,19 @@ class StreamSocket(dns._asyncbackend.StreamSocket): class Backend(dns._asyncbackend.Backend): def name(self): - return 'curio' - - async def make_socket(self, af, socktype, proto=0, - source=None, destination=None, timeout=None, - ssl_context=None, server_hostname=None): + return "curio" + + async def make_socket( + self, + af, + socktype, + proto=0, + source=None, + destination=None, + timeout=None, + ssl_context=None, + server_hostname=None, + ): if socktype == socket.SOCK_DGRAM: s = curio.socket.socket(af, socktype, proto) try: @@ -96,13 +106,17 @@ class Backend(dns._asyncbackend.Backend): else: source_addr = None async with _maybe_timeout(timeout): - s = await curio.open_connection(destination[0], destination[1], - ssl=ssl_context, - source_addr=source_addr, - server_hostname=server_hostname) + s = await curio.open_connection( + destination[0], + destination[1], + ssl=ssl_context, + source_addr=source_addr, + server_hostname=server_hostname, + ) return StreamSocket(s) - raise NotImplementedError('unsupported socket ' + - f'type {socktype}') # pragma: no cover + raise NotImplementedError( + "unsupported socket " + f"type {socktype}" + ) # pragma: no cover async def sleep(self, interval): await curio.sleep(interval) diff --git a/dns/_immutable_ctx.py b/dns/_immutable_ctx.py index ececdbe..63c0a2d 100644 --- a/dns/_immutable_ctx.py +++ b/dns/_immutable_ctx.py @@ -8,7 +8,7 @@ import contextvars import inspect -_in__init__ = contextvars.ContextVar('_immutable_in__init__', default=False) +_in__init__ = contextvars.ContextVar("_immutable_in__init__", default=False) class _Immutable: @@ -41,6 +41,7 @@ def _immutable_init(f): f(*args, **kwargs) finally: _in__init__.reset(previous) + nf.__signature__ = inspect.signature(f) return nf @@ -50,7 +51,7 @@ def immutable(cls): # Some ancestor already has the mixin, so just make sure we keep # following the __init__ protocol. cls.__init__ = _immutable_init(cls.__init__) - if hasattr(cls, '__setstate__'): + if hasattr(cls, "__setstate__"): cls.__setstate__ = _immutable_init(cls.__setstate__) ncls = cls else: @@ -63,7 +64,8 @@ def immutable(cls): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - if hasattr(cls, '__setstate__'): + if hasattr(cls, "__setstate__"): + @_immutable_init def __setstate__(self, *args, **kwargs): super().__setstate__(*args, **kwargs) diff --git a/dns/_trio_backend.py b/dns/_trio_backend.py index 8a337e9..b0c0210 100644 --- a/dns/_trio_backend.py +++ b/dns/_trio_backend.py @@ -32,7 +32,9 @@ class DatagramSocket(dns._asyncbackend.DatagramSocket): async def sendto(self, what, destination, timeout): with _maybe_timeout(timeout): return await self.socket.sendto(what, destination) - raise dns.exception.Timeout(timeout=timeout) # pragma: no cover lgtm[py/unreachable-statement] + raise dns.exception.Timeout( + timeout=timeout + ) # pragma: no cover lgtm[py/unreachable-statement] async def recvfrom(self, size, timeout): with _maybe_timeout(timeout): @@ -83,11 +85,19 @@ class StreamSocket(dns._asyncbackend.StreamSocket): class Backend(dns._asyncbackend.Backend): def name(self): - return 'trio' - - async def make_socket(self, af, socktype, proto=0, source=None, - destination=None, timeout=None, - ssl_context=None, server_hostname=None): + return "trio" + + async def make_socket( + self, + af, + socktype, + proto=0, + source=None, + destination=None, + timeout=None, + ssl_context=None, + server_hostname=None, + ): s = trio.socket.socket(af, socktype, proto) stream = None try: @@ -107,14 +117,16 @@ class Backend(dns._asyncbackend.Backend): if ssl_context: tls = True try: - stream = trio.SSLStream(stream, ssl_context, - server_hostname=server_hostname) + stream = trio.SSLStream( + stream, ssl_context, server_hostname=server_hostname + ) except Exception: # pragma: no cover await stream.aclose() raise return StreamSocket(af, stream, tls) - raise NotImplementedError('unsupported socket ' + - f'type {socktype}') # pragma: no cover + raise NotImplementedError( + "unsupported socket " + f"type {socktype}" + ) # pragma: no cover async def sleep(self, interval): await trio.sleep(interval) diff --git a/dns/asyncbackend.py b/dns/asyncbackend.py index ffd6d67..c7565a9 100644 --- a/dns/asyncbackend.py +++ b/dns/asyncbackend.py @@ -6,7 +6,12 @@ import dns.exception # pylint: disable=unused-import -from dns._asyncbackend import Socket, DatagramSocket, StreamSocket, Backend # noqa: F401 lgtm[py/unused-import] +from dns._asyncbackend import ( + Socket, + DatagramSocket, + StreamSocket, + Backend, +) # noqa: F401 lgtm[py/unused-import] # pylint: enable=unused-import @@ -17,6 +22,7 @@ _backends: Dict[str, Backend] = {} # Allow sniffio import to be disabled for testing purposes _no_sniffio = False + class AsyncLibraryNotFoundError(dns.exception.DNSException): pass @@ -33,17 +39,20 @@ def get_backend(name: str) -> Backend: backend = _backends.get(name) if backend: return backend - if name == 'trio': + if name == "trio": import dns._trio_backend + backend = dns._trio_backend.Backend() - elif name == 'curio': + elif name == "curio": import dns._curio_backend + backend = dns._curio_backend.Backend() - elif name == 'asyncio': + elif name == "asyncio": import dns._asyncio_backend + backend = dns._asyncio_backend.Backend() else: - raise NotImplementedError(f'unimplemented async backend {name}') + raise NotImplementedError(f"unimplemented async backend {name}") _backends[name] = backend return backend @@ -60,23 +69,25 @@ def sniff() -> str: if _no_sniffio: raise ImportError import sniffio + try: return sniffio.current_async_library() except sniffio.AsyncLibraryNotFoundError: - raise AsyncLibraryNotFoundError('sniffio cannot determine ' + - 'async library') + raise AsyncLibraryNotFoundError( + "sniffio cannot determine " + "async library" + ) except ImportError: import asyncio + try: asyncio.get_running_loop() - return 'asyncio' + return "asyncio" except RuntimeError: - raise AsyncLibraryNotFoundError('no async library detected') + raise AsyncLibraryNotFoundError("no async library detected") def get_default_backend() -> Backend: - """Get the default backend, initializing it if necessary. - """ + """Get the default backend, initializing it if necessary.""" if _default_backend: return _default_backend diff --git a/dns/asyncquery.py b/dns/asyncquery.py index 977f0d4..28e124d 100644 --- a/dns/asyncquery.py +++ b/dns/asyncquery.py @@ -34,8 +34,16 @@ import dns.rdataclass import dns.rdatatype import dns.transaction -from dns.query import _compute_times, _matches_destination, BadResponse, ssl, \ - UDPMode, _have_httpx, _have_http2, NoDOH +from dns.query import ( + _compute_times, + _matches_destination, + BadResponse, + ssl, + UDPMode, + _have_httpx, + _have_http2, + NoDOH, +) if _have_httpx: import httpx @@ -50,11 +58,11 @@ def _source_tuple(af, address, port): if address or port: if address is None: if af == socket.AF_INET: - address = '0.0.0.0' + address = "0.0.0.0" elif af == socket.AF_INET6: - address = '::' + address = "::" else: - raise NotImplementedError(f'unknown address family {af}') + raise NotImplementedError(f"unknown address family {af}") return (address, port) else: return None @@ -69,9 +77,12 @@ def _timeout(expiration, now=None): return None -async def send_udp(sock: dns.asyncbackend.DatagramSocket, - what: Union[dns.message.Message, bytes], destination: Any, - expiration: Optional[float]=None) -> Tuple[int, float]: +async def send_udp( + sock: dns.asyncbackend.DatagramSocket, + what: Union[dns.message.Message, bytes], + destination: Any, + expiration: Optional[float] = None, +) -> Tuple[int, float]: """Send a DNS message to the specified UDP socket. *sock*, a ``dns.asyncbackend.DatagramSocket``. @@ -95,11 +106,17 @@ async def send_udp(sock: dns.asyncbackend.DatagramSocket, return (n, sent_time) -async def receive_udp(sock: dns.asyncbackend.DatagramSocket, - destination: Optional[Any]=None, expiration: Optional[float]=None, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, - keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]]=None, request_mac: Optional[bytes]=b'', - ignore_trailing: bool=False, raise_on_truncation: bool=False) -> Any: +async def receive_udp( + sock: dns.asyncbackend.DatagramSocket, + destination: Optional[Any] = None, + expiration: Optional[float] = None, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]] = None, + request_mac: Optional[bytes] = b"", + ignore_trailing: bool = False, + raise_on_truncation: bool = False, +) -> Any: """Read a DNS message from a UDP socket. *sock*, a ``dns.asyncbackend.DatagramSocket``. @@ -108,24 +125,39 @@ async def receive_udp(sock: dns.asyncbackend.DatagramSocket, parameters, exceptions, and return type of this method. """ - wire = b'' + wire = b"" while 1: (wire, from_address) = await sock.recvfrom(65535, _timeout(expiration)) - if _matches_destination(sock.family, from_address, destination, - ignore_unexpected): + if _matches_destination( + sock.family, from_address, destination, ignore_unexpected + ): break received_time = time.time() - r = dns.message.from_wire(wire, keyring=keyring, request_mac=request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing, - raise_on_truncation=raise_on_truncation) + r = dns.message.from_wire( + wire, + keyring=keyring, + request_mac=request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + raise_on_truncation=raise_on_truncation, + ) return (r, received_time, from_address) -async def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - raise_on_truncation: bool=False, sock: Optional[dns.asyncbackend.DatagramSocket]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> dns.message.Message: + +async def udp( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + raise_on_truncation: bool = False, + sock: Optional[dns.asyncbackend.DatagramSocket] = None, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via UDP. *sock*, a ``dns.asyncbackend.DatagramSocket``, or ``None``, @@ -156,16 +188,20 @@ async def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, dtuple = (where, port) else: dtuple = None - s = await backend.make_socket(af, socket.SOCK_DGRAM, 0, stuple, - dtuple) + s = await backend.make_socket(af, socket.SOCK_DGRAM, 0, stuple, dtuple) assert s is not None await send_udp(s, wire, destination, expiration) - (r, received_time, _) = await receive_udp(s, destination, expiration, - ignore_unexpected, - one_rr_per_rrset, - q.keyring, q.mac, - ignore_trailing, - raise_on_truncation) + (r, received_time, _) = await receive_udp( + s, + destination, + expiration, + ignore_unexpected, + one_rr_per_rrset, + q.keyring, + q.mac, + ignore_trailing, + raise_on_truncation, + ) r.time = received_time - begin_time if not q.is_response(r): raise BadResponse @@ -174,12 +210,21 @@ async def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, if not sock and s: await s.close() -async def udp_with_fallback(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - udp_sock: Optional[dns.asyncbackend.DatagramSocket]=None, - tcp_sock: Optional[dns.asyncbackend.StreamSocket]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> Tuple[dns.message.Message, bool]: + +async def udp_with_fallback( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + udp_sock: Optional[dns.asyncbackend.DatagramSocket] = None, + tcp_sock: Optional[dns.asyncbackend.StreamSocket] = None, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> Tuple[dns.message.Message, bool]: """Return the response to the query, trying UDP first and falling back to TCP if UDP results in a truncated response. @@ -201,20 +246,42 @@ async def udp_with_fallback(q: dns.message.Message, where: str, timeout: Optiona method. """ try: - response = await udp(q, where, timeout, port, source, source_port, - ignore_unexpected, one_rr_per_rrset, - ignore_trailing, True, udp_sock, backend) + response = await udp( + q, + where, + timeout, + port, + source, + source_port, + ignore_unexpected, + one_rr_per_rrset, + ignore_trailing, + True, + udp_sock, + backend, + ) return (response, False) except dns.message.Truncated: - response = await tcp(q, where, timeout, port, source, source_port, - one_rr_per_rrset, ignore_trailing, tcp_sock, - backend) + response = await tcp( + q, + where, + timeout, + port, + source, + source_port, + one_rr_per_rrset, + ignore_trailing, + tcp_sock, + backend, + ) return (response, True) -async def send_tcp(sock: dns.asyncbackend.StreamSocket, - what: Union[dns.message.Message, bytes], - expiration: Optional[float]=None) -> Tuple[int, float]: +async def send_tcp( + sock: dns.asyncbackend.StreamSocket, + what: Union[dns.message.Message, bytes], + expiration: Optional[float] = None, +) -> Tuple[int, float]: """Send a DNS message to the specified TCP socket. *sock*, a ``dns.asyncbackend.StreamSocket``. @@ -241,21 +308,24 @@ async def _read_exactly(sock, count, expiration): """Read the specified number of bytes from stream. Keep trying until we either get the desired amount, or we hit EOF. """ - s = b'' + s = b"" while count > 0: n = await sock.recv(count, _timeout(expiration)) - if n == b'': + if n == b"": raise EOFError count = count - len(n) s = s + n return s -async def receive_tcp(sock: dns.asyncbackend.StreamSocket, - expiration: Optional[float]=None, one_rr_per_rrset: bool=False, - keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]]=None, - request_mac: Optional[bytes]=b'', - ignore_trailing: bool=False) -> Tuple[dns.message.Message, float]: +async def receive_tcp( + sock: dns.asyncbackend.StreamSocket, + expiration: Optional[float] = None, + one_rr_per_rrset: bool = False, + keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]] = None, + request_mac: Optional[bytes] = b"", + ignore_trailing: bool = False, +) -> Tuple[dns.message.Message, float]: """Read a DNS message from a TCP socket. *sock*, a ``dns.asyncbackend.StreamSocket``. @@ -268,17 +338,28 @@ async def receive_tcp(sock: dns.asyncbackend.StreamSocket, (l,) = struct.unpack("!H", ldata) wire = await _read_exactly(sock, l, expiration) received_time = time.time() - r = dns.message.from_wire(wire, keyring=keyring, request_mac=request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing) + r = dns.message.from_wire( + wire, + keyring=keyring, + request_mac=request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + ) return (r, received_time) -async def tcp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - sock: Optional[dns.asyncbackend.StreamSocket]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> dns.message.Message: +async def tcp( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + sock: Optional[dns.asyncbackend.StreamSocket] = None, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via TCP. *sock*, a ``dns.asyncbacket.StreamSocket``, or ``None``, the @@ -313,13 +394,14 @@ async def tcp(q: dns.message.Message, where: str, timeout: Optional[float]=None, dtuple = (where, port) if not backend: backend = dns.asyncbackend.get_default_backend() - s = await backend.make_socket(af, socket.SOCK_STREAM, 0, stuple, - dtuple, timeout) + s = await backend.make_socket( + af, socket.SOCK_STREAM, 0, stuple, dtuple, timeout + ) assert s is not None await send_tcp(s, wire, expiration) - (r, received_time) = await receive_tcp(s, expiration, one_rr_per_rrset, - q.keyring, q.mac, - ignore_trailing) + (r, received_time) = await receive_tcp( + s, expiration, one_rr_per_rrset, q.keyring, q.mac, ignore_trailing + ) r.time = received_time - begin_time if not q.is_response(r): raise BadResponse @@ -328,13 +410,21 @@ async def tcp(q: dns.message.Message, where: str, timeout: Optional[float]=None, if not sock and s: await s.close() -async def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, - port: int=853, source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - sock: Optional[dns.asyncbackend.StreamSocket]=None, - backend: Optional[dns.asyncbackend.Backend]=None, - ssl_context: Optional[ssl.SSLContext]=None, - server_hostname: Optional[str]=None) -> dns.message.Message: + +async def tls( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 853, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + sock: Optional[dns.asyncbackend.StreamSocket] = None, + backend: Optional[dns.asyncbackend.Backend] = None, + ssl_context: Optional[ssl.SSLContext] = None, + server_hostname: Optional[str] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via TLS. *sock*, an ``asyncbackend.StreamSocket``, or ``None``, the socket @@ -367,15 +457,32 @@ async def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, dtuple = (where, port) if not backend: backend = dns.asyncbackend.get_default_backend() - s = await backend.make_socket(af, socket.SOCK_STREAM, 0, stuple, - dtuple, timeout, ssl_context, - server_hostname) + s = await backend.make_socket( + af, + socket.SOCK_STREAM, + 0, + stuple, + dtuple, + timeout, + ssl_context, + server_hostname, + ) else: s = sock try: timeout = _timeout(expiration) - response = await tcp(q, where, timeout, port, source, source_port, - one_rr_per_rrset, ignore_trailing, s, backend) + response = await tcp( + q, + where, + timeout, + port, + source, + source_port, + one_rr_per_rrset, + ignore_trailing, + s, + backend, + ) end_time = time.time() response.time = end_time - begin_time return response @@ -383,11 +490,21 @@ async def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, if not sock and s: await s.close() -async def https(q: dns.message.Message, where: str, timeout: Optional[float]=None, - port: int=443, source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - client: Optional[httpx.AsyncClient]=None, - path: str='/dns-query', post: bool=True, verify: bool=True) -> dns.message.Message: + +async def https( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 443, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + client: Optional[httpx.AsyncClient] = None, + path: str = "/dns-query", + post: bool = True, + verify: bool = True, +) -> dns.message.Message: """Return the response obtained after sending a query via DNS-over-HTTPS. *client*, a ``httpx.AsyncClient``. If provided, the client to use for @@ -401,7 +518,7 @@ async def https(q: dns.message.Message, where: str, timeout: Optional[float]=Non """ if not _have_httpx: - raise NoDOH('httpx is not available.') # pragma: no cover + raise NoDOH("httpx is not available.") # pragma: no cover wire = q.to_wire() try: @@ -409,14 +526,12 @@ async def https(q: dns.message.Message, where: str, timeout: Optional[float]=Non except ValueError: af = None transport = None - headers = { - "accept": "application/dns-message" - } + headers = {"accept": "application/dns-message"} if af is not None: if af == socket.AF_INET: - url = 'https://{}:{}{}'.format(where, port, path) + url = "https://{}:{}{}".format(where, port, path) elif af == socket.AF_INET6: - url = 'https://[{}]:{}{}'.format(where, port, path) + url = "https://[{}]:{}{}".format(where, port, path) else: url = where if source is not None: @@ -426,24 +541,29 @@ async def https(q: dns.message.Message, where: str, timeout: Optional[float]=Non client_to_close = None try: if not client: - client = httpx.AsyncClient(http1=True, http2=_have_http2, - verify=verify, transport=transport) + client = httpx.AsyncClient( + http1=True, http2=_have_http2, verify=verify, transport=transport + ) client_to_close = client # see https://tools.ietf.org/html/rfc8484#section-4.1.1 for DoH # GET and POST examples if post: - headers.update({ - "content-type": "application/dns-message", - "content-length": str(len(wire)) - }) - response = await client.post(url, headers=headers, content=wire, - timeout=timeout) + headers.update( + { + "content-type": "application/dns-message", + "content-length": str(len(wire)), + } + ) + response = await client.post( + url, headers=headers, content=wire, timeout=timeout + ) else: wire = base64.urlsafe_b64encode(wire).rstrip(b"=") twire = wire.decode() # httpx does a repr() if we give it bytes - response = await client.get(url, headers=headers, timeout=timeout, - params={"dns": twire}) + response = await client.get( + url, headers=headers, timeout=timeout, params={"dns": twire} + ) finally: if client_to_close: await client_to_close.aclose() @@ -451,25 +571,37 @@ async def https(q: dns.message.Message, where: str, timeout: Optional[float]=Non # see https://tools.ietf.org/html/rfc8484#section-4.2.1 for info about DoH # status codes if response.status_code < 200 or response.status_code > 299: - raise ValueError('{} responded with status code {}' - '\nResponse body: {!r}'.format(where, - response.status_code, - response.content)) - r = dns.message.from_wire(response.content, - keyring=q.keyring, - request_mac=q.request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing) + raise ValueError( + "{} responded with status code {}" + "\nResponse body: {!r}".format( + where, response.status_code, response.content + ) + ) + r = dns.message.from_wire( + response.content, + keyring=q.keyring, + request_mac=q.request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + ) r.time = response.elapsed.total_seconds() if not q.is_response(r): raise BadResponse return r -async def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManager, - query: Optional[dns.message.Message]=None, - port: int=53, timeout: Optional[float]=None, lifetime: Optional[float]=None, - source: Optional[str]=None, source_port: int=0, udp_mode: UDPMode=UDPMode.NEVER, - backend: Optional[dns.asyncbackend.Backend]=None) -> None: + +async def inbound_xfr( + where: str, + txn_manager: dns.transaction.TransactionManager, + query: Optional[dns.message.Message] = None, + port: int = 53, + timeout: Optional[float] = None, + lifetime: Optional[float] = None, + source: Optional[str] = None, + source_port: int = 0, + udp_mode: UDPMode = UDPMode.NEVER, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> None: """Conduct an inbound transfer and apply it via a transaction from the txn_manager. @@ -502,42 +634,48 @@ async def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManage is_udp = False if not backend: backend = dns.asyncbackend.get_default_backend() - s = await backend.make_socket(af, sock_type, 0, stuple, dtuple, - _timeout(expiration)) + s = await backend.make_socket( + af, sock_type, 0, stuple, dtuple, _timeout(expiration) + ) async with s: if is_udp: await s.sendto(wire, dtuple, _timeout(expiration)) else: tcpmsg = struct.pack("!H", len(wire)) + wire await s.sendall(tcpmsg, expiration) - with dns.xfr.Inbound(txn_manager, rdtype, serial, - is_udp) as inbound: + with dns.xfr.Inbound(txn_manager, rdtype, serial, is_udp) as inbound: done = False tsig_ctx = None while not done: (_, mexpiration) = _compute_times(timeout) - if mexpiration is None or \ - (expiration is not None and mexpiration > expiration): + if mexpiration is None or ( + expiration is not None and mexpiration > expiration + ): mexpiration = expiration if is_udp: destination = _lltuple((where, port), af) while True: timeout = _timeout(mexpiration) - (rwire, from_address) = await s.recvfrom(65535, - timeout) - if _matches_destination(af, from_address, - destination, True): + (rwire, from_address) = await s.recvfrom(65535, timeout) + if _matches_destination( + af, from_address, destination, True + ): break else: ldata = await _read_exactly(s, 2, mexpiration) (l,) = struct.unpack("!H", ldata) rwire = await _read_exactly(s, l, mexpiration) - is_ixfr = (rdtype == dns.rdatatype.IXFR) - r = dns.message.from_wire(rwire, keyring=query.keyring, - request_mac=query.mac, xfr=True, - origin=origin, tsig_ctx=tsig_ctx, - multi=(not is_udp), - one_rr_per_rrset=is_ixfr) + is_ixfr = rdtype == dns.rdatatype.IXFR + r = dns.message.from_wire( + rwire, + keyring=query.keyring, + request_mac=query.mac, + xfr=True, + origin=origin, + tsig_ctx=tsig_ctx, + multi=(not is_udp), + one_rr_per_rrset=is_ixfr, + ) try: done = inbound.process_message(r) except dns.xfr.UseTCP: diff --git a/dns/asyncresolver.py b/dns/asyncresolver.py index e196dbb..14a25a3 100644 --- a/dns/asyncresolver.py +++ b/dns/asyncresolver.py @@ -42,13 +42,19 @@ _tcp = dns.asyncquery.tcp class Resolver(dns.resolver.BaseResolver): """Asynchronous DNS stub resolver.""" - async def resolve(self, qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - tcp: bool=False, source: Optional[str]=None, - raise_on_no_answer: bool=True, source_port: int=0, - lifetime: Optional[float]=None, search: Optional[bool]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> dns.resolver.Answer: + async def resolve( + self, + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, + search: Optional[bool] = None, + backend: Optional[dns.asyncbackend.Backend] = None, + ) -> dns.resolver.Answer: """Query nameservers asynchronously to find the answer to the question. *backend*, a ``dns.asyncbackend.Backend``, or ``None``. If ``None``, @@ -59,8 +65,9 @@ class Resolver(dns.resolver.BaseResolver): type of this method. """ - resolution = dns.resolver._Resolution(self, qname, rdtype, rdclass, tcp, - raise_on_no_answer, search) + resolution = dns.resolver._Resolution( + self, qname, rdtype, rdclass, tcp, raise_on_no_answer, search + ) if not backend: backend = dns.asyncbackend.get_default_backend() start = time.time() @@ -79,25 +86,34 @@ class Resolver(dns.resolver.BaseResolver): (nameserver, port, tcp, backoff) = resolution.next_nameserver() if backoff: await backend.sleep(backoff) - timeout = self._compute_timeout(start, lifetime, - resolution.errors) + timeout = self._compute_timeout(start, lifetime, resolution.errors) try: if dns.inet.is_address(nameserver): if tcp: - response = await _tcp(request, nameserver, - timeout, port, - source, source_port, - backend=backend) + response = await _tcp( + request, + nameserver, + timeout, + port, + source, + source_port, + backend=backend, + ) else: - response = await _udp(request, nameserver, - timeout, port, - source, source_port, - raise_on_truncation=True, - backend=backend) + response = await _udp( + request, + nameserver, + timeout, + port, + source, + source_port, + raise_on_truncation=True, + backend=backend, + ) else: - response = await dns.asyncquery.https(request, - nameserver, - timeout=timeout) + response = await dns.asyncquery.https( + request, nameserver, timeout=timeout + ) except Exception as ex: (_, done) = resolution.query_result(None, ex) continue @@ -109,7 +125,9 @@ class Resolver(dns.resolver.BaseResolver): if answer is not None: return answer - async def resolve_address(self, ipaddr: str, *args: Any, **kwargs: Dict[str, Any]) -> dns.resolver.Answer: + async def resolve_address( + self, ipaddr: str, *args: Any, **kwargs: Dict[str, Any] + ) -> dns.resolver.Answer: """Use an asynchronous resolver to run a reverse query for PTR records. @@ -129,10 +147,11 @@ class Resolver(dns.resolver.BaseResolver): # in the kwargs more than once. modified_kwargs: Dict[str, Any] = {} modified_kwargs.update(kwargs) - modified_kwargs['rdtype'] = dns.rdatatype.PTR - modified_kwargs['rdclass'] = dns.rdataclass.IN - return await self.resolve(dns.reversename.from_address(ipaddr), - *args, **modified_kwargs) + modified_kwargs["rdtype"] = dns.rdatatype.PTR + modified_kwargs["rdclass"] = dns.rdataclass.IN + return await self.resolve( + dns.reversename.from_address(ipaddr), *args, **modified_kwargs + ) # pylint: disable=redefined-outer-name @@ -180,13 +199,18 @@ def reset_default_resolver() -> None: default_resolver = Resolver() -async def resolve(qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - tcp: bool=False, source: Optional[str]=None, - raise_on_no_answer: bool=True, source_port: int=0, - lifetime: Optional[float]=None, search: Optional[bool]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> dns.resolver.Answer: +async def resolve( + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, + search: Optional[bool] = None, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> dns.resolver.Answer: """Query nameservers asynchronously to find the answer to the question. This is a convenience function that uses the default resolver @@ -196,13 +220,23 @@ async def resolve(qname: Union[dns.name.Name, str], information on the parameters. """ - return await get_default_resolver().resolve(qname, rdtype, rdclass, tcp, - source, raise_on_no_answer, - source_port, lifetime, search, - backend) - - -async def resolve_address(ipaddr: str, *args: Any, **kwargs: Dict[str, Any]) -> dns.resolver.Answer: + return await get_default_resolver().resolve( + qname, + rdtype, + rdclass, + tcp, + source, + raise_on_no_answer, + source_port, + lifetime, + search, + backend, + ) + + +async def resolve_address( + ipaddr: str, *args: Any, **kwargs: Dict[str, Any] +) -> dns.resolver.Answer: """Use a resolver to run a reverse query for PTR records. See :py:func:`dns.asyncresolver.Resolver.resolve_address` for more @@ -211,6 +245,7 @@ async def resolve_address(ipaddr: str, *args: Any, **kwargs: Dict[str, Any]) -> return await get_default_resolver().resolve_address(ipaddr, *args, **kwargs) + async def canonical_name(name: Union[dns.name.Name, str]) -> dns.name.Name: """Determine the canonical name of *name*. @@ -220,10 +255,14 @@ async def canonical_name(name: Union[dns.name.Name, str]) -> dns.name.Name: return await get_default_resolver().canonical_name(name) -async def zone_for_name(name: Union[dns.name.Name, str], - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, - tcp: bool=False, resolver: Optional[Resolver]=None, - backend: Optional[dns.asyncbackend.Backend]=None) -> dns.name.Name: + +async def zone_for_name( + name: Union[dns.name.Name, str], + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + tcp: bool = False, + resolver: Optional[Resolver] = None, + backend: Optional[dns.asyncbackend.Backend] = None, +) -> dns.name.Name: """Find the name of the zone which contains the specified name. See :py:func:`dns.resolver.Resolver.zone_for_name` for more @@ -238,8 +277,9 @@ async def zone_for_name(name: Union[dns.name.Name, str], raise NotAbsolute(name) while True: try: - answer = await resolver.resolve(name, dns.rdatatype.SOA, rdclass, - tcp, backend=backend) + answer = await resolver.resolve( + name, dns.rdatatype.SOA, rdclass, tcp, backend=backend + ) assert answer.rrset is not None if answer.rrset.name == name: return name diff --git a/dns/dnssec.py b/dns/dnssec.py index 331f4af..b325f9f 100644 --- a/dns/dnssec.py +++ b/dns/dnssec.py @@ -83,17 +83,19 @@ def key_id(key: DNSKEY) -> int: else: total = 0 for i in range(len(rdata) // 2): - total += (rdata[2 * i] << 8) + \ - rdata[2 * i + 1] + total += (rdata[2 * i] << 8) + rdata[2 * i + 1] if len(rdata) % 2 != 0: total += rdata[len(rdata) - 1] << 8 - total += ((total >> 16) & 0xffff) - return total & 0xffff + total += (total >> 16) & 0xFFFF + return total & 0xFFFF -def make_ds(name: Union[dns.name.Name, str], key: dns.rdata.Rdata, - algorithm: Union[DSDigest, str], - origin: Optional[dns.name.Name]=None) -> DS: +def make_ds( + name: Union[dns.name.Name, str], + key: dns.rdata.Rdata, + algorithm: Union[DSDigest, str], + origin: Optional[dns.name.Name] = None, +) -> DS: """Create a DS record for a DNSSEC key. *name*, a ``dns.name.Name`` or ``str``, the owner name of the DS record. @@ -118,7 +120,7 @@ def make_ds(name: Union[dns.name.Name, str], key: dns.rdata.Rdata, except Exception: raise UnsupportedAlgorithm('unsupported algorithm "%s"' % algorithm) if not isinstance(key, DNSKEY): - raise ValueError('key is not a DNSKEY') + raise ValueError("key is not a DNSKEY") if algorithm == DSDigest.SHA1: dshash = hashlib.sha1() elif algorithm == DSDigest.SHA256: @@ -136,15 +138,16 @@ def make_ds(name: Union[dns.name.Name, str], key: dns.rdata.Rdata, dshash.update(key.to_wire(origin=origin)) digest = dshash.digest() - dsrdata = struct.pack("!HBB", key_id(key), key.algorithm, algorithm) + \ - digest - ds = dns.rdata.from_wire(dns.rdataclass.IN, dns.rdatatype.DS, dsrdata, 0, - len(dsrdata)) + dsrdata = struct.pack("!HBB", key_id(key), key.algorithm, algorithm) + digest + ds = dns.rdata.from_wire( + dns.rdataclass.IN, dns.rdatatype.DS, dsrdata, 0, len(dsrdata) + ) return cast(DS, ds) -def _find_candidate_keys(keys: Dict[dns.name.Name, Union[dns.rdataset.Rdataset, dns.node.Node]], - rrsig: RRSIG) -> Optional[List[DNSKEY]]: +def _find_candidate_keys( + keys: Dict[dns.name.Name, Union[dns.rdataset.Rdataset, dns.node.Node]], rrsig: RRSIG +) -> Optional[List[DNSKEY]]: value = keys.get(rrsig.signer) if isinstance(value, dns.node.Node): rdataset = value.get_rdataset(dns.rdataclass.IN, dns.rdatatype.DNSKEY) @@ -152,14 +155,21 @@ def _find_candidate_keys(keys: Dict[dns.name.Name, Union[dns.rdataset.Rdataset, rdataset = value if rdataset is None: return None - return [cast(DNSKEY, rd) for rd in rdataset if - rd.algorithm == rrsig.algorithm and key_id(rd) == rrsig.key_tag] + return [ + cast(DNSKEY, rd) + for rd in rdataset + if rd.algorithm == rrsig.algorithm and key_id(rd) == rrsig.key_tag + ] def _is_rsa(algorithm: int) -> bool: - return algorithm in (Algorithm.RSAMD5, Algorithm.RSASHA1, - Algorithm.RSASHA1NSEC3SHA1, Algorithm.RSASHA256, - Algorithm.RSASHA512) + return algorithm in ( + Algorithm.RSAMD5, + Algorithm.RSASHA1, + Algorithm.RSASHA1NSEC3SHA1, + Algorithm.RSASHA256, + Algorithm.RSASHA512, + ) def _is_dsa(algorithm: int) -> bool: @@ -183,8 +193,12 @@ def _is_md5(algorithm: int) -> bool: def _is_sha1(algorithm: int) -> bool: - return algorithm in (Algorithm.DSA, Algorithm.RSASHA1, - Algorithm.DSANSEC3SHA1, Algorithm.RSASHA1NSEC3SHA1) + return algorithm in ( + Algorithm.DSA, + Algorithm.RSASHA1, + Algorithm.DSANSEC3SHA1, + Algorithm.RSASHA1NSEC3SHA1, + ) def _is_sha256(algorithm: int) -> bool: @@ -215,35 +229,36 @@ def _make_hash(algorithm: int) -> Any: if algorithm == Algorithm.ED448: return hashes.SHAKE256(114) - raise ValidationFailure('unknown hash for algorithm %u' % algorithm) + raise ValidationFailure("unknown hash for algorithm %u" % algorithm) def _bytes_to_long(b: bytes) -> int: - return int.from_bytes(b, 'big') + return int.from_bytes(b, "big") def _validate_signature(sig: bytes, data: bytes, key: DNSKEY, chosen_hash: Any) -> None: keyptr: bytes if _is_rsa(key.algorithm): - # we ignore because mypy is confused and thinks key.key is a str for unknown reasons. + # we ignore because mypy is confused and thinks key.key is a str for unknown + # reasons. keyptr = key.key - (bytes_,) = struct.unpack('!B', keyptr[0:1]) + (bytes_,) = struct.unpack("!B", keyptr[0:1]) keyptr = keyptr[1:] if bytes_ == 0: - (bytes_,) = struct.unpack('!H', keyptr[0:2]) + (bytes_,) = struct.unpack("!H", keyptr[0:2]) keyptr = keyptr[2:] rsa_e = keyptr[0:bytes_] rsa_n = keyptr[bytes_:] try: rsa_public_key = rsa.RSAPublicNumbers( - _bytes_to_long(rsa_e), - _bytes_to_long(rsa_n)).public_key(default_backend()) + _bytes_to_long(rsa_e), _bytes_to_long(rsa_n) + ).public_key(default_backend()) except ValueError: - raise ValidationFailure('invalid public key') + raise ValidationFailure("invalid public key") rsa_public_key.verify(sig, data, padding.PKCS1v15(), chosen_hash) elif _is_dsa(key.algorithm): keyptr = key.key - (t,) = struct.unpack('!B', keyptr[0:1]) + (t,) = struct.unpack("!B", keyptr[0:1]) keyptr = keyptr[1:] octets = 64 + t * 8 dsa_q = keyptr[0:20] @@ -257,11 +272,11 @@ def _validate_signature(sig: bytes, data: bytes, key: DNSKEY, chosen_hash: Any) dsa_public_key = dsa.DSAPublicNumbers( _bytes_to_long(dsa_y), dsa.DSAParameterNumbers( - _bytes_to_long(dsa_p), - _bytes_to_long(dsa_q), - _bytes_to_long(dsa_g))).public_key(default_backend()) + _bytes_to_long(dsa_p), _bytes_to_long(dsa_q), _bytes_to_long(dsa_g) + ), + ).public_key(default_backend()) except ValueError: - raise ValidationFailure('invalid public key') + raise ValidationFailure("invalid public key") dsa_public_key.verify(sig, data, chosen_hash) elif _is_ecdsa(key.algorithm): keyptr = key.key @@ -273,14 +288,13 @@ def _validate_signature(sig: bytes, data: bytes, key: DNSKEY, chosen_hash: Any) curve = ec.SECP384R1() octets = 48 ecdsa_x = keyptr[0:octets] - ecdsa_y = keyptr[octets:octets * 2] + ecdsa_y = keyptr[octets : octets * 2] try: ecdsa_public_key = ec.EllipticCurvePublicNumbers( - curve=curve, - x=_bytes_to_long(ecdsa_x), - y=_bytes_to_long(ecdsa_y)).public_key(default_backend()) + curve=curve, x=_bytes_to_long(ecdsa_x), y=_bytes_to_long(ecdsa_y) + ).public_key(default_backend()) except ValueError: - raise ValidationFailure('invalid public key') + raise ValidationFailure("invalid public key") ecdsa_public_key.verify(sig, data, ec.ECDSA(chosen_hash)) elif _is_eddsa(key.algorithm): keyptr = key.key @@ -292,20 +306,24 @@ def _validate_signature(sig: bytes, data: bytes, key: DNSKEY, chosen_hash: Any) try: eddsa_public_key = loader.from_public_bytes(keyptr) except ValueError: - raise ValidationFailure('invalid public key') + raise ValidationFailure("invalid public key") eddsa_public_key.verify(sig, data) elif _is_gost(key.algorithm): raise UnsupportedAlgorithm( - 'algorithm "%s" not supported by dnspython' % - algorithm_to_text(key.algorithm)) + 'algorithm "%s" not supported by dnspython' + % algorithm_to_text(key.algorithm) + ) else: - raise ValidationFailure('unknown algorithm %u' % key.algorithm) + raise ValidationFailure("unknown algorithm %u" % key.algorithm) -def _validate_rrsig(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], - rrsig: RRSIG, - keys: Dict[dns.name.Name, Union[dns.node.Node, dns.rdataset.Rdataset]], - origin: Optional[dns.name.Name]=None, now: Optional[float]=None) -> None: +def _validate_rrsig( + rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], + rrsig: RRSIG, + keys: Dict[dns.name.Name, Union[dns.node.Node, dns.rdataset.Rdataset]], + origin: Optional[dns.name.Name] = None, + now: Optional[float] = None, +) -> None: """Validate an RRset against a single signature rdata, throwing an exception if validation is not successful. @@ -340,7 +358,7 @@ def _validate_rrsig(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdata candidate_keys = _find_candidate_keys(keys, rrsig) if candidate_keys is None: - raise ValidationFailure('unknown key') + raise ValidationFailure("unknown key") # For convenience, allow the rrset to be specified as a (name, # rdataset) tuple as well as a proper rrset @@ -354,15 +372,14 @@ def _validate_rrsig(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdata if now is None: now = time.time() if rrsig.expiration < now: - raise ValidationFailure('expired') + raise ValidationFailure("expired") if rrsig.inception > now: - raise ValidationFailure('not yet valid') + raise ValidationFailure("not yet valid") if _is_dsa(rrsig.algorithm): sig_r = rrsig.signature[1:21] sig_s = rrsig.signature[21:] - sig = utils.encode_dss_signature(_bytes_to_long(sig_r), - _bytes_to_long(sig_s)) + sig = utils.encode_dss_signature(_bytes_to_long(sig_r), _bytes_to_long(sig_s)) elif _is_ecdsa(rrsig.algorithm): if rrsig.algorithm == Algorithm.ECDSAP256SHA256: octets = 32 @@ -370,34 +387,32 @@ def _validate_rrsig(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdata octets = 48 sig_r = rrsig.signature[0:octets] sig_s = rrsig.signature[octets:] - sig = utils.encode_dss_signature(_bytes_to_long(sig_r), - _bytes_to_long(sig_s)) + sig = utils.encode_dss_signature(_bytes_to_long(sig_r), _bytes_to_long(sig_s)) else: sig = rrsig.signature - data = b'' + data = b"" data += rrsig.to_wire(origin=origin)[:18] data += rrsig.signer.to_digestable(origin) # Derelativize the name before considering labels. if not rrname.is_absolute(): if origin is None: - raise ValidationFailure('relative RR name without an origin specified') + raise ValidationFailure("relative RR name without an origin specified") rrname = rrname.derelativize(origin) if len(rrname) - 1 < rrsig.labels: - raise ValidationFailure('owner name longer than RRSIG labels') + raise ValidationFailure("owner name longer than RRSIG labels") elif rrsig.labels < len(rrname) - 1: suffix = rrname.split(rrsig.labels + 1)[1] - rrname = dns.name.from_text('*', suffix) + rrname = dns.name.from_text("*", suffix) rrnamebuf = rrname.to_digestable() - rrfixed = struct.pack('!HHI', rdataset.rdtype, rdataset.rdclass, - rrsig.original_ttl) + rrfixed = struct.pack("!HHI", rdataset.rdtype, rdataset.rdclass, rrsig.original_ttl) rdatas = [rdata.to_digestable(origin) for rdata in rdataset] for rdata in sorted(rdatas): data += rrnamebuf data += rrfixed - rrlen = struct.pack('!H', len(rdata)) + rrlen = struct.pack("!H", len(rdata)) data += rrlen data += rdata @@ -411,13 +426,16 @@ def _validate_rrsig(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdata # this happens on an individual validation failure continue # nothing verified -- raise failure: - raise ValidationFailure('verify failure') + raise ValidationFailure("verify failure") -def _validate(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], - rrsigset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], - keys: Dict[dns.name.Name, Union[dns.node.Node, dns.rdataset.Rdataset]], - origin: Optional[dns.name.Name]=None, now: Optional[float]=None) -> None: +def _validate( + rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], + rrsigset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rdataset]], + keys: Dict[dns.name.Name, Union[dns.node.Node, dns.rdataset.Rdataset]], + origin: Optional[dns.name.Name] = None, + now: Optional[float] = None, +) -> None: """Validate an RRset against a signature RRset, throwing an exception if none of the signatures validate. @@ -468,7 +486,7 @@ def _validate(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rd for rrsig in rrsigrdataset: if not isinstance(rrsig, RRSIG): - raise ValidationFailure('expected an RRSIG') + raise ValidationFailure("expected an RRSIG") try: _validate_rrsig(rrset, rrsig, keys, origin, now) return @@ -477,8 +495,12 @@ def _validate(rrset: Union[dns.rrset.RRset, Tuple[dns.name.Name, dns.rdataset.Rd raise ValidationFailure("no RRSIGs validated") -def nsec3_hash(domain: Union[dns.name.Name, str], salt: Optional[Union[str, bytes]], - iterations: int, algorithm: Union[int, str]) -> str: +def nsec3_hash( + domain: Union[dns.name.Name, str], + salt: Optional[Union[str, bytes]], + iterations: int, + algorithm: Union[int, str], +) -> str: """ Calculate the NSEC3 hash, according to https://tools.ietf.org/html/rfc5155#section-5 @@ -510,7 +532,7 @@ def nsec3_hash(domain: Union[dns.name.Name, str], salt: Optional[Union[str, byte raise ValueError("Wrong hash algorithm (only SHA1 is supported)") if salt is None: - salt_encoded = b'' + salt_encoded = b"" elif isinstance(salt, str): if len(salt) % 2 == 0: salt_encoded = bytes.fromhex(salt) @@ -535,8 +557,9 @@ def nsec3_hash(domain: Union[dns.name.Name, str], salt: Optional[Union[str, byte def _need_pyca(*args, **kwargs): - raise ImportError("DNSSEC validation requires " + - "python cryptography") # pragma: no cover + raise ImportError( + "DNSSEC validation requires " + "python cryptography" + ) # pragma: no cover try: @@ -555,8 +578,8 @@ except ImportError: # pragma: no cover validate_rrsig = _need_pyca _have_pyca = False else: - validate = _validate # type: ignore - validate_rrsig = _validate_rrsig # type: ignore + validate = _validate # type: ignore + validate_rrsig = _validate_rrsig # type: ignore _have_pyca = True ### BEGIN generated Algorithm constants diff --git a/dns/e164.py b/dns/e164.py index 6e34ae5..453736d 100644 --- a/dns/e164.py +++ b/dns/e164.py @@ -24,10 +24,12 @@ import dns.name import dns.resolver #: The public E.164 domain. -public_enum_domain = dns.name.from_text('e164.arpa.') +public_enum_domain = dns.name.from_text("e164.arpa.") -def from_e164(text: str, origin: Optional[dns.name.Name]=public_enum_domain) -> dns.name.Name: +def from_e164( + text: str, origin: Optional[dns.name.Name] = public_enum_domain +) -> dns.name.Name: """Convert an E.164 number in textual form into a Name object whose value is the ENUM domain name for that number. @@ -44,11 +46,14 @@ def from_e164(text: str, origin: Optional[dns.name.Name]=public_enum_domain) -> parts = [d for d in text if d.isdigit()] parts.reverse() - return dns.name.from_text('.'.join(parts), origin=origin) + return dns.name.from_text(".".join(parts), origin=origin) -def to_e164(name: dns.name.Name, origin: Optional[dns.name.Name]=public_enum_domain, - want_plus_prefix: bool=True) -> str: +def to_e164( + name: dns.name.Name, + origin: Optional[dns.name.Name] = public_enum_domain, + want_plus_prefix: bool = True, +) -> str: """Convert an ENUM domain name into an E.164 number. Note that dnspython does not have any information about preferred @@ -72,16 +77,19 @@ def to_e164(name: dns.name.Name, origin: Optional[dns.name.Name]=public_enum_dom name = name.relativize(origin) dlabels = [d for d in name.labels if d.isdigit() and len(d) == 1] if len(dlabels) != len(name.labels): - raise dns.exception.SyntaxError('non-digit labels in ENUM domain name') + raise dns.exception.SyntaxError("non-digit labels in ENUM domain name") dlabels.reverse() - text = b''.join(dlabels) + text = b"".join(dlabels) if want_plus_prefix: - text = b'+' + text + text = b"+" + text return text.decode() -def query(number: str, domains: Iterable[Union[dns.name.Name, str]], - resolver: Optional[dns.resolver.Resolver]=None) -> dns.resolver.Answer: +def query( + number: str, + domains: Iterable[Union[dns.name.Name, str]], + resolver: Optional[dns.resolver.Resolver] = None, +) -> dns.resolver.Answer: """Look for NAPTR RRs for the specified number in the specified domains. e.g. lookup('16505551212', ['e164.dnspython.org.', 'e164.arpa.']) @@ -102,7 +110,7 @@ def query(number: str, domains: Iterable[Union[dns.name.Name, str]], domain = dns.name.from_text(domain) qname = dns.e164.from_e164(number, domain) try: - return resolver.resolve(qname, 'NAPTR') + return resolver.resolve(qname, "NAPTR") except dns.resolver.NXDOMAIN as e: e_nx += e raise e_nx diff --git a/dns/edns.py b/dns/edns.py index d4dca55..64436cd 100644 --- a/dns/edns.py +++ b/dns/edns.py @@ -69,7 +69,7 @@ class Option: """ self.otype = OptionType.make(otype) - def to_wire(self, file: Optional[Any]=None) -> Optional[bytes]: + def to_wire(self, file: Optional[Any] = None) -> Optional[bytes]: """Convert an option to wire format. Returns a ``bytes`` or ``None``. @@ -78,7 +78,7 @@ class Option: raise NotImplementedError # pragma: no cover @classmethod - def from_wire_parser(cls, otype: OptionType, parser: 'dns.wire.Parser') -> 'Option': + def from_wire_parser(cls, otype: OptionType, parser: "dns.wire.Parser") -> "Option": """Build an EDNS option object from wire format. *otype*, a ``dns.edns.OptionType``, is the option type. @@ -118,26 +118,22 @@ class Option: return self._cmp(other) != 0 def __lt__(self, other): - if not isinstance(other, Option) or \ - self.otype != other.otype: + if not isinstance(other, Option) or self.otype != other.otype: return NotImplemented return self._cmp(other) < 0 def __le__(self, other): - if not isinstance(other, Option) or \ - self.otype != other.otype: + if not isinstance(other, Option) or self.otype != other.otype: return NotImplemented return self._cmp(other) <= 0 def __ge__(self, other): - if not isinstance(other, Option) or \ - self.otype != other.otype: + if not isinstance(other, Option) or self.otype != other.otype: return NotImplemented return self._cmp(other) >= 0 def __gt__(self, other): - if not isinstance(other, Option) or \ - self.otype != other.otype: + if not isinstance(other, Option) or self.otype != other.otype: return NotImplemented return self._cmp(other) > 0 @@ -157,7 +153,7 @@ class GenericOption(Option): # lgtm[py/missing-equals] super().__init__(otype) self.data = dns.rdata.Rdata._as_bytes(data, True) - def to_wire(self, file: Optional[Any]=None) -> Optional[bytes]: + def to_wire(self, file: Optional[Any] = None) -> Optional[bytes]: if file: file.write(self.data) return None @@ -168,14 +164,16 @@ class GenericOption(Option): # lgtm[py/missing-equals] return "Generic %d" % self.otype @classmethod - def from_wire_parser(cls, otype: Union[OptionType, str], parser: 'dns.wire.Parser') -> Option: + def from_wire_parser( + cls, otype: Union[OptionType, str], parser: "dns.wire.Parser" + ) -> Option: return cls(otype, parser.get_remaining()) class ECSOption(Option): # lgtm[py/missing-equals] """EDNS Client Subnet (ECS, RFC7871)""" - def __init__(self, address: str, srclen: Optional[int]=None, scopelen: int=0): + def __init__(self, address: str, srclen: Optional[int] = None, scopelen: int = 0): """*address*, a ``str``, is the client address information. *srclen*, an ``int``, the source prefix length, which is the @@ -204,7 +202,7 @@ class ECSOption(Option): # lgtm[py/missing-equals] srclen = dns.rdata.Rdata._as_int(srclen, 0, 32) scopelen = dns.rdata.Rdata._as_int(scopelen, 0, 32) else: # pragma: no cover (this will never happen) - raise ValueError('Bad address family') + raise ValueError("Bad address family") assert srclen is not None self.address = address @@ -219,13 +217,11 @@ class ECSOption(Option): # lgtm[py/missing-equals] self.addrdata = addrdata[:nbytes] nbits = srclen % 8 if nbits != 0: - last = struct.pack('B', - ord(self.addrdata[-1:]) & (0xff << (8 - nbits))) + last = struct.pack("B", ord(self.addrdata[-1:]) & (0xFF << (8 - nbits))) self.addrdata = self.addrdata[:-1] + last def to_text(self) -> str: - return "ECS {}/{} scope/{}".format(self.address, self.srclen, - self.scopelen) + return "ECS {}/{} scope/{}".format(self.address, self.srclen, self.scopelen) @staticmethod def from_text(text: str) -> Option: @@ -251,7 +247,7 @@ class ECSOption(Option): # lgtm[py/missing-equals] >>> # it understands results from `dns.edns.ECSOption.to_text()` >>> dns.edns.ECSOption.from_text('ECS 1.2.3.4/24/32') """ - optional_prefix = 'ECS' + optional_prefix = "ECS" tokens = text.split() ecs_text = None if len(tokens) == 1: @@ -262,29 +258,32 @@ class ECSOption(Option): # lgtm[py/missing-equals] ecs_text = tokens[1] else: raise ValueError('could not parse ECS from "{}"'.format(text)) - n_slashes = ecs_text.count('/') + n_slashes = ecs_text.count("/") if n_slashes == 1: - address, tsrclen = ecs_text.split('/') - tscope = '0' + address, tsrclen = ecs_text.split("/") + tscope = "0" elif n_slashes == 2: - address, tsrclen, tscope = ecs_text.split('/') + address, tsrclen, tscope = ecs_text.split("/") else: raise ValueError('could not parse ECS from "{}"'.format(text)) try: scope = int(tscope) except ValueError: - raise ValueError('invalid scope ' + - '"{}": scope must be an integer'.format(tscope)) + raise ValueError( + "invalid scope " + '"{}": scope must be an integer'.format(tscope) + ) try: srclen = int(tsrclen) except ValueError: - raise ValueError('invalid srclen ' + - '"{}": srclen must be an integer'.format(tsrclen)) + raise ValueError( + "invalid srclen " + '"{}": srclen must be an integer'.format(tsrclen) + ) return ECSOption(address, srclen, scope) - def to_wire(self, file: Optional[Any]=None) -> Optional[bytes]: - value = (struct.pack('!HBB', self.family, self.srclen, self.scopelen) + - self.addrdata) + def to_wire(self, file: Optional[Any] = None) -> Optional[bytes]: + value = ( + struct.pack("!HBB", self.family, self.srclen, self.scopelen) + self.addrdata + ) if file: file.write(value) return None @@ -292,18 +291,20 @@ class ECSOption(Option): # lgtm[py/missing-equals] return value @classmethod - def from_wire_parser(cls, otype: Union[OptionType, str], parser: 'dns.wire.Parser') -> Option: - family, src, scope = parser.get_struct('!HBB') + def from_wire_parser( + cls, otype: Union[OptionType, str], parser: "dns.wire.Parser" + ) -> Option: + family, src, scope = parser.get_struct("!HBB") addrlen = int(math.ceil(src / 8.0)) prefix = parser.get_bytes(addrlen) if family == 1: pad = 4 - addrlen - addr = dns.ipv4.inet_ntoa(prefix + b'\x00' * pad) + addr = dns.ipv4.inet_ntoa(prefix + b"\x00" * pad) elif family == 2: pad = 16 - addrlen - addr = dns.ipv6.inet_ntoa(prefix + b'\x00' * pad) + addr = dns.ipv6.inet_ntoa(prefix + b"\x00" * pad) else: - raise ValueError('unsupported family') + raise ValueError("unsupported family") return cls(addr, src, scope) @@ -343,7 +344,7 @@ class EDECode(dns.enum.IntEnum): class EDEOption(Option): # lgtm[py/missing-equals] """Extended DNS Error (EDE, RFC8914)""" - def __init__(self, code: Union[EDECode, str], text: Optional[str]=None): + def __init__(self, code: Union[EDECode, str], text: Optional[str] = None): """*code*, a ``dns.edns.EDECode`` or ``str``, the info code of the extended error. @@ -355,19 +356,19 @@ class EDEOption(Option): # lgtm[py/missing-equals] self.code = EDECode.make(code) if text is not None and not isinstance(text, str): - raise ValueError('text must be string or None') + raise ValueError("text must be string or None") self.text = text def to_text(self) -> str: - output = f'EDE {self.code}' + output = f"EDE {self.code}" if self.text is not None: - output += f': {self.text}' + output += f": {self.text}" return output - def to_wire(self, file: Optional[Any]=None) -> Optional[bytes]: - value = struct.pack('!H', self.code) + def to_wire(self, file: Optional[Any] = None) -> Optional[bytes]: + value = struct.pack("!H", self.code) if self.text is not None: - value += self.text.encode('utf8') + value += self.text.encode("utf8") if file: file.write(value) @@ -376,14 +377,16 @@ class EDEOption(Option): # lgtm[py/missing-equals] return value @classmethod - def from_wire_parser(cls, otype: Union[OptionType, str], parser: 'dns.wire.Parser') -> Option: + def from_wire_parser( + cls, otype: Union[OptionType, str], parser: "dns.wire.Parser" + ) -> Option: the_code = EDECode.make(parser.get_uint16()) text = parser.get_remaining() if text: if text[-1] == 0: # text MAY be null-terminated text = text[:-1] - btext = text.decode('utf8') + btext = text.decode("utf8") else: btext = None @@ -409,7 +412,9 @@ def get_option_class(otype: OptionType) -> Any: return cls -def option_from_wire_parser(otype: Union[OptionType, str], parser: 'dns.wire.Parser') -> Option: +def option_from_wire_parser( + otype: Union[OptionType, str], parser: "dns.wire.Parser" +) -> Option: """Build an EDNS option object from wire format. *otype*, an ``int``, is the option type. @@ -424,7 +429,9 @@ def option_from_wire_parser(otype: Union[OptionType, str], parser: 'dns.wire.Par return cls.from_wire_parser(otype, parser) -def option_from_wire(otype: Union[OptionType, str], wire: bytes, current: int, olen: int) -> Option: +def option_from_wire( + otype: Union[OptionType, str], wire: bytes, current: int, olen: int +) -> Option: """Build an EDNS option object from wire format. *otype*, an ``int``, is the option type. @@ -442,6 +449,7 @@ def option_from_wire(otype: Union[OptionType, str], wire: bytes, current: int, o with parser.restrict_to(olen): return option_from_wire_parser(otype, parser) + def register_type(implementation: Any, otype: OptionType) -> None: """Register the implementation of an option type. @@ -452,6 +460,7 @@ def register_type(implementation: Any, otype: OptionType) -> None: _type_to_class[otype] = implementation + ### BEGIN generated OptionType constants NSID = OptionType.NSID diff --git a/dns/entropy.py b/dns/entropy.py index 7da2e04..5010356 100644 --- a/dns/entropy.py +++ b/dns/entropy.py @@ -21,10 +21,11 @@ import os import hashlib import random import time + try: import threading as _threading except ImportError: # pragma: no cover - import dummy_threading as _threading # type: ignore + import dummy_threading as _threading # type: ignore class EntropyPool: @@ -34,14 +35,14 @@ class EntropyPool: # leaving this code doesn't hurt anything as the library code # is used if present. - def __init__(self, seed: Optional[bytes]=None): + def __init__(self, seed: Optional[bytes] = None): self.pool_index = 0 self.digest: Optional[bytearray] = None self.next_byte = 0 self.lock = _threading.Lock() self.hash = hashlib.sha1() self.hash_len = 20 - self.pool = bytearray(b'\0' * self.hash_len) + self.pool = bytearray(b"\0" * self.hash_len) if seed is not None: self._stir(seed) self.seeded = True @@ -54,7 +55,7 @@ class EntropyPool: for c in entropy: if self.pool_index == self.hash_len: self.pool_index = 0 - b = c & 0xff + b = c & 0xFF self.pool[self.pool_index] ^= b self.pool_index += 1 @@ -68,7 +69,7 @@ class EntropyPool: seed = os.urandom(16) except Exception: # pragma: no cover try: - with open('/dev/urandom', 'rb', 0) as r: + with open("/dev/urandom", "rb", 0) as r: seed = r.read(16) except Exception: seed = str(time.time()).encode() @@ -99,7 +100,7 @@ class EntropyPool: def random_between(self, first: int, last: int) -> int: size = last - first + 1 if size > 4294967296: - raise ValueError('too big') + raise ValueError("too big") if size > 65536: rand = self.random_32 max = 4294967295 @@ -111,6 +112,7 @@ class EntropyPool: max = 255 return first + size * rand() // (max + 1) + pool = EntropyPool() system_random: Optional[Any] @@ -119,12 +121,14 @@ try: except Exception: # pragma: no cover system_random = None + def random_16() -> int: if system_random is not None: return system_random.randrange(0, 65536) else: return pool.random_16() + def between(first: int, last: int) -> int: if system_random is not None: return system_random.randrange(first, last + 1) diff --git a/dns/enum.py b/dns/enum.py index b822dd5..9c67488 100644 --- a/dns/enum.py +++ b/dns/enum.py @@ -17,6 +17,7 @@ import enum + class IntEnum(enum.IntEnum): @classmethod def _check_value(cls, value): @@ -33,8 +34,8 @@ class IntEnum(enum.IntEnum): except KeyError: pass prefix = cls._prefix() - if text.startswith(prefix) and text[len(prefix):].isdigit(): - value = int(text[len(prefix):]) + if text.startswith(prefix) and text[len(prefix) :].isdigit(): + value = int(text[len(prefix) :]) cls._check_value(value) try: return cls(value) @@ -83,7 +84,7 @@ class IntEnum(enum.IntEnum): @classmethod def _prefix(cls): - return '' + return "" @classmethod def _unknown_exception_class(cls): diff --git a/dns/exception.py b/dns/exception.py index aa0144d..3b2f1cd 100644 --- a/dns/exception.py +++ b/dns/exception.py @@ -73,14 +73,15 @@ class DNSException(Exception): For sanity we do not allow to mix old and new behavior.""" if args or kwargs: - assert bool(args) != bool(kwargs), \ - 'keyword arguments are mutually exclusive with positional args' + assert bool(args) != bool( + kwargs + ), "keyword arguments are mutually exclusive with positional args" def _check_kwargs(self, **kwargs): if kwargs: - assert set(kwargs.keys()) == self.supp_kwargs, \ - 'following set of keyword args is required: %s' % ( - self.supp_kwargs) + assert ( + set(kwargs.keys()) == self.supp_kwargs + ), "following set of keyword args is required: %s" % (self.supp_kwargs) return kwargs def _fmt_kwargs(self, **kwargs): @@ -129,10 +130,12 @@ class TooBig(DNSException): class Timeout(DNSException): """The DNS operation timed out.""" - supp_kwargs = {'timeout'} + + supp_kwargs = {"timeout"} fmt = "The DNS operation timed out after {timeout:.3f} seconds" - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -145,7 +148,6 @@ class ExceptionWrapper: return self def __exit__(self, exc_type, exc_val, exc_tb): - if exc_type is not None and not isinstance(exc_val, - self.exception_class): + if exc_type is not None and not isinstance(exc_val, self.exception_class): raise self.exception_class(str(exc_val)) from exc_val return False diff --git a/dns/flags.py b/dns/flags.py index 6fe1afd..b21b8e3 100644 --- a/dns/flags.py +++ b/dns/flags.py @@ -23,6 +23,7 @@ import enum # Standard DNS flags + class Flag(enum.IntFlag): #: Query Response QR = 0x8000 @@ -42,6 +43,7 @@ class Flag(enum.IntFlag): # EDNS flags + class EDNSFlag(enum.IntFlag): #: DNSSEC answer OK DO = 0x8000 @@ -60,7 +62,7 @@ def _to_text(flags: int, enum_class: Any) -> str: for k, v in enum_class.__members__.items(): if flags & v != 0: text_flags.append(k) - return ' '.join(text_flags) + return " ".join(text_flags) def from_text(text: str) -> int: @@ -102,6 +104,7 @@ def edns_to_text(flags: int) -> str: return _to_text(flags, EDNSFlag) + ### BEGIN generated Flag constants QR = Flag.QR diff --git a/dns/grange.py b/dns/grange.py index ebb64d2..3a52278 100644 --- a/dns/grange.py +++ b/dns/grange.py @@ -21,6 +21,7 @@ from typing import Tuple import dns + def from_text(text: str) -> Tuple[int, int, int]: """Convert the text form of a range in a ``$GENERATE`` statement to an integer. @@ -33,22 +34,22 @@ def from_text(text: str) -> Tuple[int, int, int]: start = -1 stop = -1 step = 1 - cur = '' + cur = "" state = 0 # state 0 1 2 # x - y / z - if text and text[0] == '-': + if text and text[0] == "-": raise dns.exception.SyntaxError("Start cannot be a negative number") for c in text: - if c == '-' and state == 0: + if c == "-" and state == 0: start = int(cur) - cur = '' + cur = "" state = 1 - elif c == '/': + elif c == "/": stop = int(cur) - cur = '' + cur = "" state = 2 elif c.isdigit(): cur += c @@ -66,6 +67,6 @@ def from_text(text: str) -> Tuple[int, int, int]: assert step >= 1 assert start >= 0 if start > stop: - raise dns.exception.SyntaxError('start must be <= stop') + raise dns.exception.SyntaxError("start must be <= stop") return (start, stop, step) diff --git a/dns/immutable.py b/dns/immutable.py index 8a42621..38fbe59 100644 --- a/dns/immutable.py +++ b/dns/immutable.py @@ -9,7 +9,7 @@ from dns._immutable_ctx import immutable @immutable class Dict(collections.abc.Mapping): # lgtm[py/missing-equals] - def __init__(self, dictionary: Any, no_copy: bool=False): + def __init__(self, dictionary: Any, no_copy: bool = False): """Make an immutable dictionary from the specified dictionary. If *no_copy* is `True`, then *dictionary* will be wrapped instead @@ -30,7 +30,7 @@ class Dict(collections.abc.Mapping): # lgtm[py/missing-equals] h = 0 for key in sorted(self._odict.keys()): h ^= hash(key) - object.__setattr__(self, '_hash', h) + object.__setattr__(self, "_hash", h) # this does return an int, but pylint doesn't figure that out return self._hash diff --git a/dns/inet.py b/dns/inet.py index b3ed999..11180c9 100644 --- a/dns/inet.py +++ b/dns/inet.py @@ -137,7 +137,9 @@ def is_address(text: str) -> bool: return False -def low_level_address_tuple(high_tuple: Tuple[str, int], af: Optional[int]=None) -> Any: +def low_level_address_tuple( + high_tuple: Tuple[str, int], af: Optional[int] = None +) -> Any: """Given a "high-level" address tuple, i.e. an (address, port) return the appropriate "low-level" address tuple suitable for use in socket calls. @@ -152,13 +154,13 @@ def low_level_address_tuple(high_tuple: Tuple[str, int], af: Optional[int]=None) if af == AF_INET: return (address, port) elif af == AF_INET6: - i = address.find('%') + i = address.find("%") if i < 0: # no scope, shortcut! return (address, port, 0, 0) # try to avoid getaddrinfo() addrpart = address[:i] - scope = address[i + 1:] + scope = address[i + 1 :] if scope.isdigit(): return (addrpart, port, 0, int(scope)) try: @@ -168,4 +170,4 @@ def low_level_address_tuple(high_tuple: Tuple[str, int], af: Optional[int]=None) ((*_, tup), *_) = socket.getaddrinfo(address, port, flags=ai_flags) return tup else: - raise NotImplementedError(f'unknown address family {af}') + raise NotImplementedError(f"unknown address family {af}") diff --git a/dns/ipv4.py b/dns/ipv4.py index fddad1b..b8e148f 100644 --- a/dns/ipv4.py +++ b/dns/ipv4.py @@ -23,6 +23,7 @@ import struct import dns.exception + def inet_ntoa(address: bytes) -> str: """Convert an IPv4 address in binary form to text form. @@ -33,8 +34,8 @@ def inet_ntoa(address: bytes) -> str: if len(address) != 4: raise dns.exception.SyntaxError - return ('%u.%u.%u.%u' % (address[0], address[1], - address[2], address[3])) + return "%u.%u.%u.%u" % (address[0], address[1], address[2], address[3]) + def inet_aton(text: Union[str, bytes]) -> bytes: """Convert an IPv4 address in text form to binary form. @@ -48,17 +49,17 @@ def inet_aton(text: Union[str, bytes]) -> bytes: btext = text.encode() else: btext = text - parts = btext.split(b'.') + parts = btext.split(b".") if len(parts) != 4: raise dns.exception.SyntaxError for part in parts: if not part.isdigit(): raise dns.exception.SyntaxError - if len(part) > 1 and part[0] == ord('0'): + if len(part) > 1 and part[0] == ord("0"): # No leading zeros raise dns.exception.SyntaxError try: b = [int(part) for part in parts] - return struct.pack('BBBB', *b) + return struct.pack("BBBB", *b) except Exception: raise dns.exception.SyntaxError diff --git a/dns/ipv6.py b/dns/ipv6.py index 9e6e8b6..fbd4962 100644 --- a/dns/ipv6.py +++ b/dns/ipv6.py @@ -25,7 +25,8 @@ import binascii import dns.exception import dns.ipv4 -_leading_zero = re.compile(r'0+([0-9a-f]+)') +_leading_zero = re.compile(r"0+([0-9a-f]+)") + def inet_ntoa(address: bytes) -> str: """Convert an IPv6 address in binary form to text form. @@ -43,7 +44,7 @@ def inet_ntoa(address: bytes) -> str: i = 0 l = len(hex) while i < l: - chunk = hex[i:i + 4].decode() + chunk = hex[i : i + 4].decode() # strip leading zeros. we do this with an re instead of # with lstrip() because lstrip() didn't support chars until # python 2.2.2 @@ -60,7 +61,7 @@ def inet_ntoa(address: bytes) -> str: start = -1 last_was_zero = False for i in range(8): - if chunks[i] != '0': + if chunks[i] != "0": if last_was_zero: end = i current_len = end - start @@ -78,27 +79,30 @@ def inet_ntoa(address: bytes) -> str: best_start = start best_len = current_len if best_len > 1: - if best_start == 0 and \ - (best_len == 6 or - best_len == 5 and chunks[5] == 'ffff'): + if best_start == 0 and (best_len == 6 or best_len == 5 and chunks[5] == "ffff"): # We have an embedded IPv4 address if best_len == 6: - prefix = '::' + prefix = "::" else: - prefix = '::ffff:' + prefix = "::ffff:" thex = prefix + dns.ipv4.inet_ntoa(address[12:]) else: - thex = ':'.join(chunks[:best_start]) + '::' + \ - ':'.join(chunks[best_start + best_len:]) + thex = ( + ":".join(chunks[:best_start]) + + "::" + + ":".join(chunks[best_start + best_len :]) + ) else: - thex = ':'.join(chunks) + thex = ":".join(chunks) return thex -_v4_ending = re.compile(br'(.*):(\d+\.\d+\.\d+\.\d+)$') -_colon_colon_start = re.compile(br'::.*') -_colon_colon_end = re.compile(br'.*::$') -def inet_aton(text: Union[str, bytes], ignore_scope: bool=False) -> bytes: +_v4_ending = re.compile(rb"(.*):(\d+\.\d+\.\d+\.\d+)$") +_colon_colon_start = re.compile(rb"::.*") +_colon_colon_end = re.compile(rb".*::$") + + +def inet_aton(text: Union[str, bytes], ignore_scope: bool = False) -> bytes: """Convert an IPv6 address in text form to binary form. *text*, a ``str``, the IPv6 address in textual form. @@ -118,30 +122,32 @@ def inet_aton(text: Union[str, bytes], ignore_scope: bool=False) -> bytes: btext = text if ignore_scope: - parts = btext.split(b'%') + parts = btext.split(b"%") l = len(parts) if l == 2: btext = parts[0] elif l > 2: raise dns.exception.SyntaxError - if btext == b'': + if btext == b"": raise dns.exception.SyntaxError - elif btext.endswith(b':') and not btext.endswith(b'::'): + elif btext.endswith(b":") and not btext.endswith(b"::"): raise dns.exception.SyntaxError - elif btext.startswith(b':') and not btext.startswith(b'::'): + elif btext.startswith(b":") and not btext.startswith(b"::"): raise dns.exception.SyntaxError - elif btext == b'::': - btext = b'0::' + elif btext == b"::": + btext = b"0::" # # Get rid of the icky dot-quad syntax if we have it. # m = _v4_ending.match(btext) if m is not None: b = dns.ipv4.inet_aton(m.group(2)) - btext = ("{}:{:02x}{:02x}:{:02x}{:02x}".format(m.group(1).decode(), - b[0], b[1], b[2], - b[3])).encode() + btext = ( + "{}:{:02x}{:02x}:{:02x}{:02x}".format( + m.group(1).decode(), b[0], b[1], b[2], b[3] + ) + ).encode() # # Try to turn '::<whatever>' into ':<whatever>'; if no match try to # turn '<whatever>::' into '<whatever>:' @@ -156,29 +162,29 @@ def inet_aton(text: Union[str, bytes], ignore_scope: bool=False) -> bytes: # # Now canonicalize into 8 chunks of 4 hex digits each # - chunks = btext.split(b':') + chunks = btext.split(b":") l = len(chunks) if l > 8: raise dns.exception.SyntaxError seen_empty = False canonical: List[bytes] = [] for c in chunks: - if c == b'': + if c == b"": if seen_empty: raise dns.exception.SyntaxError seen_empty = True for _ in range(0, 8 - l + 1): - canonical.append(b'0000') + canonical.append(b"0000") else: lc = len(c) if lc > 4: raise dns.exception.SyntaxError if lc != 4: - c = (b'0' * (4 - lc)) + c + c = (b"0" * (4 - lc)) + c canonical.append(c) if l < 8 and not seen_empty: raise dns.exception.SyntaxError - btext = b''.join(canonical) + btext = b"".join(canonical) # # Finally we can go to binary. @@ -188,7 +194,9 @@ def inet_aton(text: Union[str, bytes], ignore_scope: bool=False) -> bytes: except (binascii.Error, TypeError): raise dns.exception.SyntaxError -_mapped_prefix = b'\x00' * 10 + b'\xff\xff' + +_mapped_prefix = b"\x00" * 10 + b"\xff\xff" + def is_mapped(address: bytes) -> bool: """Is the specified address a mapped IPv4 address? diff --git a/dns/message.py b/dns/message.py index 0e1e433..967fefe 100644 --- a/dns/message.py +++ b/dns/message.py @@ -73,9 +73,10 @@ class UnknownTSIGKey(dns.exception.DNSException): class Truncated(dns.exception.DNSException): """The truncated flag is set.""" - supp_kwargs = {'message'} + supp_kwargs = {"message"} - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -84,7 +85,7 @@ class Truncated(dns.exception.DNSException): Returns a ``dns.message.Message``. """ - return self.kwargs['message'] + return self.kwargs["message"] class NotQueryResponse(dns.exception.DNSException): @@ -98,12 +99,14 @@ class ChainTooLong(dns.exception.DNSException): class AnswerForNXDOMAIN(dns.exception.DNSException): """The rcode is NXDOMAIN but an answer was found.""" + class NoPreviousName(dns.exception.SyntaxError): """No previous name was known.""" class MessageSection(dns.enum.IntEnum): """Message sections""" + QUESTION = 0 ANSWER = 1 AUTHORITY = 2 @@ -123,18 +126,24 @@ class MessageError: DEFAULT_EDNS_PAYLOAD = 1232 MAX_CHAIN = 16 -IndexKeyType = Tuple[int, dns.name.Name, dns.rdataclass.RdataClass, - dns.rdatatype.RdataType, Optional[dns.rdatatype.RdataType], - Optional[dns.rdataclass.RdataClass]] +IndexKeyType = Tuple[ + int, + dns.name.Name, + dns.rdataclass.RdataClass, + dns.rdatatype.RdataType, + Optional[dns.rdatatype.RdataType], + Optional[dns.rdataclass.RdataClass], +] IndexType = Dict[IndexKeyType, dns.rrset.RRset] SectionType = Union[int, List[dns.rrset.RRset]] + class Message: """A DNS message.""" _section_enum = MessageSection - def __init__(self, id: Optional[int]=None): + def __init__(self, id: Optional[int] = None): if id is None: self.id = dns.entropy.random_16() else: @@ -145,7 +154,7 @@ class Message: self.request_payload = 0 self.keyring: Any = None self.tsig: Optional[dns.rrset.RRset] = None - self.request_mac = b'' + self.request_mac = b"" self.xfr = False self.origin: Optional[dns.name.Name] = None self.tsig_ctx: Optional[Any] = None @@ -155,7 +164,7 @@ class Message: @property def question(self) -> List[dns.rrset.RRset]: - """ The question section.""" + """The question section.""" return self.sections[0] @question.setter @@ -164,7 +173,7 @@ class Message: @property def answer(self) -> List[dns.rrset.RRset]: - """ The answer section.""" + """The answer section.""" return self.sections[1] @answer.setter @@ -173,7 +182,7 @@ class Message: @property def authority(self) -> List[dns.rrset.RRset]: - """ The authority section.""" + """The authority section.""" return self.sections[2] @authority.setter @@ -182,7 +191,7 @@ class Message: @property def additional(self) -> List[dns.rrset.RRset]: - """ The additional data section.""" + """The additional data section.""" return self.sections[3] @additional.setter @@ -190,13 +199,17 @@ class Message: self.sections[3] = v def __repr__(self): - return '<DNS message, ID ' + repr(self.id) + '>' + return "<DNS message, ID " + repr(self.id) + ">" def __str__(self): return self.to_text() - def to_text(self, origin: Optional[dns.name.Name]=None, relativize: bool=True, - **kw: Dict[str, Any]) -> str: + def to_text( + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + **kw: Dict[str, Any], + ) -> str: """Convert the message to text. The *origin*, *relativize*, and any other keyword @@ -206,23 +219,22 @@ class Message: """ s = io.StringIO() - s.write('id %d\n' % self.id) - s.write('opcode %s\n' % dns.opcode.to_text(self.opcode())) - s.write('rcode %s\n' % dns.rcode.to_text(self.rcode())) - s.write('flags %s\n' % dns.flags.to_text(self.flags)) + s.write("id %d\n" % self.id) + s.write("opcode %s\n" % dns.opcode.to_text(self.opcode())) + s.write("rcode %s\n" % dns.rcode.to_text(self.rcode())) + s.write("flags %s\n" % dns.flags.to_text(self.flags)) if self.edns >= 0: - s.write('edns %s\n' % self.edns) + s.write("edns %s\n" % self.edns) if self.ednsflags != 0: - s.write('eflags %s\n' % - dns.flags.edns_to_text(self.ednsflags)) - s.write('payload %d\n' % self.payload) + s.write("eflags %s\n" % dns.flags.edns_to_text(self.ednsflags)) + s.write("payload %d\n" % self.payload) for opt in self.options: - s.write('option %s\n' % opt.to_text()) + s.write("option %s\n" % opt.to_text()) for (name, which) in self._section_enum.__members__.items(): - s.write(f';{name}\n') + s.write(f";{name}\n") for rrset in self.section_from_number(which): s.write(rrset.to_text(origin, relativize, **kw)) - s.write('\n') + s.write("\n") # # We strip off the final \n so the caller can print the result without # doing weird things to get around eccentricities in Python print @@ -256,20 +268,25 @@ class Message: def __ne__(self, other): return not self.__eq__(other) - def is_response(self, other: 'Message') -> bool: + def is_response(self, other: "Message") -> bool: """Is *other*, also a ``dns.message.Message``, a response to this message? Returns a ``bool``. """ - if other.flags & dns.flags.QR == 0 or \ - self.id != other.id or \ - dns.opcode.from_flags(self.flags) != \ - dns.opcode.from_flags(other.flags): + if ( + other.flags & dns.flags.QR == 0 + or self.id != other.id + or dns.opcode.from_flags(self.flags) != dns.opcode.from_flags(other.flags) + ): return False - if other.rcode() in {dns.rcode.FORMERR, dns.rcode.SERVFAIL, - dns.rcode.NOTIMP, dns.rcode.REFUSED}: + if other.rcode() in { + dns.rcode.FORMERR, + dns.rcode.SERVFAIL, + dns.rcode.NOTIMP, + dns.rcode.REFUSED, + }: # We don't check the question section in these cases if # the other question section is empty, even though they # still really ought to have a question section. @@ -303,7 +320,7 @@ class Message: for i, our_section in enumerate(self.sections): if section is our_section: return self._section_enum(i) - raise ValueError('unknown section') + raise ValueError("unknown section") def section_from_number(self, number: int) -> List[dns.rrset.RRset]: """Return the section list associated with the specified section @@ -320,15 +337,17 @@ class Message: section = self._section_enum.make(number) return self.sections[section] - def find_rrset(self, - section: SectionType, - name: dns.name.Name, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - deleting: Optional[dns.rdataclass.RdataClass]=None, - create: bool=False, - force_unique: bool=False) -> dns.rrset.RRset: + def find_rrset( + self, + section: SectionType, + name: dns.name.Name, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + deleting: Optional[dns.rdataclass.RdataClass] = None, + create: bool = False, + force_unique: bool = False, + ) -> dns.rrset.RRset: """Find the RRset with the given attributes in the specified section. *section*, an ``int`` section number, or one of the section @@ -378,8 +397,7 @@ class Message: return rrset else: for rrset in the_section: - if rrset.full_match(name, rdclass, rdtype, covers, - deleting): + if rrset.full_match(name, rdclass, rdtype, covers, deleting): return rrset if not create: raise KeyError @@ -389,15 +407,17 @@ class Message: self.index[key] = rrset return rrset - def get_rrset(self, - section: SectionType, - name: dns.name.Name, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - deleting: Optional[dns.rdataclass.RdataClass]=None, - create: bool=False, - force_unique: bool=False) -> Optional[dns.rrset.RRset]: + def get_rrset( + self, + section: SectionType, + name: dns.name.Name, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + deleting: Optional[dns.rdataclass.RdataClass] = None, + create: bool = False, + force_unique: bool = False, + ) -> Optional[dns.rrset.RRset]: """Get the RRset with the given attributes in the specified section. If the RRset is not found, None is returned. @@ -433,14 +453,21 @@ class Message: """ try: - rrset = self.find_rrset(section, name, rdclass, rdtype, covers, - deleting, create, force_unique) + rrset = self.find_rrset( + section, name, rdclass, rdtype, covers, deleting, create, force_unique + ) except KeyError: rrset = None return rrset - def to_wire(self, origin: Optional[dns.name.Name]=None, max_size: int=0, - multi: bool=False, tsig_ctx: Optional[Any]=None, **kw: Dict[str, Any]) -> bytes: + def to_wire( + self, + origin: Optional[dns.name.Name] = None, + max_size: int = 0, + multi: bool = False, + tsig_ctx: Optional[Any] = None, + **kw: Dict[str, Any], + ) -> bytes: """Return a string containing the message in DNS compressed wire format. @@ -490,13 +517,15 @@ class Message: r.add_rrset(dns.renderer.ADDITIONAL, rrset, **kw) r.write_header() if self.tsig is not None: - (new_tsig, ctx) = dns.tsig.sign(r.get_wire(), - self.keyring, - self.tsig[0], - int(time.time()), - self.request_mac, - tsig_ctx, - multi) + (new_tsig, ctx) = dns.tsig.sign( + r.get_wire(), + self.keyring, + self.tsig[0], + int(time.time()), + self.request_mac, + tsig_ctx, + multi, + ) self.tsig.clear() self.tsig.add(new_tsig) r.add_rrset(dns.renderer.ADDITIONAL, self.tsig) @@ -506,17 +535,32 @@ class Message: return r.get_wire() @staticmethod - def _make_tsig(keyname, algorithm, time_signed, fudge, mac, original_id, - error, other): - tsig = dns.rdtypes.ANY.TSIG.TSIG(dns.rdataclass.ANY, dns.rdatatype.TSIG, - algorithm, time_signed, fudge, mac, - original_id, error, other) + def _make_tsig( + keyname, algorithm, time_signed, fudge, mac, original_id, error, other + ): + tsig = dns.rdtypes.ANY.TSIG.TSIG( + dns.rdataclass.ANY, + dns.rdatatype.TSIG, + algorithm, + time_signed, + fudge, + mac, + original_id, + error, + other, + ) return dns.rrset.from_rdata(keyname, 0, tsig) - def use_tsig(self, keyring: Any, keyname: Optional[Union[dns.name.Name, str]]=None, - fudge: int=300, original_id: Optional[int]=None, tsig_error: int=0, - other_data: bytes=b'', - algorithm: Union[dns.name.Name, str]=dns.tsig.default_algorithm) -> None: + def use_tsig( + self, + keyring: Any, + keyname: Optional[Union[dns.name.Name, str]] = None, + fudge: int = 300, + original_id: Optional[int] = None, + tsig_error: int = 0, + other_data: bytes = b"", + algorithm: Union[dns.name.Name, str] = dns.tsig.default_algorithm, + ) -> None: """When sending, a TSIG signature using the specified key should be added. @@ -570,8 +614,16 @@ class Message: self.keyring = key if original_id is None: original_id = self.id - self.tsig = self._make_tsig(keyname, self.keyring.algorithm, 0, fudge, - b'', original_id, tsig_error, other_data) + self.tsig = self._make_tsig( + keyname, + self.keyring.algorithm, + 0, + fudge, + b"", + original_id, + tsig_error, + other_data, + ) @property def keyname(self) -> Optional[dns.name.Name]: @@ -607,13 +659,17 @@ class Message: @staticmethod def _make_opt(flags=0, payload=DEFAULT_EDNS_PAYLOAD, options=None): - opt = dns.rdtypes.ANY.OPT.OPT(payload, dns.rdatatype.OPT, - options or ()) + opt = dns.rdtypes.ANY.OPT.OPT(payload, dns.rdatatype.OPT, options or ()) return dns.rrset.from_rdata(dns.name.root, int(flags), opt) - def use_edns(self, edns: Optional[Union[int, bool]]=0, ednsflags: int=0, payload: int=DEFAULT_EDNS_PAYLOAD, - request_payload: Optional[int]=None, - options: Optional[List[dns.edns.Option]]=None) -> None: + def use_edns( + self, + edns: Optional[Union[int, bool]] = 0, + ednsflags: int = 0, + payload: int = DEFAULT_EDNS_PAYLOAD, + request_payload: Optional[int] = None, + options: Optional[List[dns.edns.Option]] = None, + ) -> None: """Configure EDNS behavior. *edns*, an ``int``, is the EDNS level to use. Specifying @@ -645,7 +701,7 @@ class Message: else: # make sure the EDNS version in ednsflags agrees with edns ednsflags &= 0xFF00FFFF - ednsflags |= (edns << 16) + ednsflags |= edns << 16 if options is None: options = [] self.opt = self._make_opt(ednsflags, payload, options) @@ -656,7 +712,7 @@ class Message: @property def edns(self) -> int: if self.opt: - return (self.ednsflags & 0xff0000) >> 16 + return (self.ednsflags & 0xFF0000) >> 16 else: return -1 @@ -688,7 +744,7 @@ class Message: else: return () - def want_dnssec(self, wanted: bool=True) -> None: + def want_dnssec(self, wanted: bool = True) -> None: """Enable or disable 'DNSSEC desired' flag in requests. *wanted*, a ``bool``. If ``True``, then DNSSEC data is @@ -746,16 +802,20 @@ class Message: # pylint: enable=unused-argument - def _parse_special_rr_header(self, section, count, position, - name, rdclass, rdtype): + def _parse_special_rr_header(self, section, count, position, name, rdclass, rdtype): if rdtype == dns.rdatatype.OPT: - if section != MessageSection.ADDITIONAL or self.opt or \ - name != dns.name.root: + if ( + section != MessageSection.ADDITIONAL + or self.opt + or name != dns.name.root + ): raise BadEDNS elif rdtype == dns.rdatatype.TSIG: - if section != MessageSection.ADDITIONAL or \ - rdclass != dns.rdatatype.ANY or \ - position != count - 1: + if ( + section != MessageSection.ADDITIONAL + or rdclass != dns.rdatatype.ANY + or position != count - 1 + ): raise BadTSIG return (rdclass, rdtype, None, False) @@ -778,8 +838,14 @@ class ChainingResult: The ``cnames`` attribute is a list of all the CNAME RRSets followed to get to the canonical name. """ - def __init__(self, canonical_name: dns.name.Name, answer: Optional[dns.rrset.RRset], - minimum_ttl: int, cnames: List[dns.rrset.RRset]): + + def __init__( + self, + canonical_name: dns.name.Name, + answer: Optional[dns.rrset.RRset], + minimum_ttl: int, + cnames: List[dns.rrset.RRset], + ): self.canonical_name = canonical_name self.answer = answer self.minimum_ttl = minimum_ttl @@ -815,16 +881,17 @@ class QueryMessage(Message): cnames = [] while count < MAX_CHAIN: try: - answer = self.find_rrset(self.answer, qname, question.rdclass, - question.rdtype) + answer = self.find_rrset( + self.answer, qname, question.rdclass, question.rdtype + ) min_ttl = min(min_ttl, answer.ttl) break except KeyError: if question.rdtype != dns.rdatatype.CNAME: try: - crrset = self.find_rrset(self.answer, qname, - question.rdclass, - dns.rdatatype.CNAME) + crrset = self.find_rrset( + self.answer, qname, question.rdclass, dns.rdatatype.CNAME + ) cnames.append(crrset) min_ttl = min(min_ttl, crrset.ttl) for rd in crrset: @@ -849,9 +916,9 @@ class QueryMessage(Message): # Look for an SOA RR whose owner name is a superdomain # of qname. try: - srrset = self.find_rrset(self.authority, auname, - question.rdclass, - dns.rdatatype.SOA) + srrset = self.find_rrset( + self.authority, auname, question.rdclass, dns.rdatatype.SOA + ) min_ttl = min(min_ttl, srrset.ttl, srrset[0].minimum) break except KeyError: @@ -915,9 +982,17 @@ class _WireReader: raising them. """ - def __init__(self, wire, initialize_message, question_only=False, - one_rr_per_rrset=False, ignore_trailing=False, - keyring=None, multi=False, continue_on_error=False): + def __init__( + self, + wire, + initialize_message, + question_only=False, + one_rr_per_rrset=False, + ignore_trailing=False, + keyring=None, + multi=False, + continue_on_error=False, + ): self.parser = dns.wire.Parser(wire) self.message = None self.initialize_message = initialize_message @@ -937,12 +1012,13 @@ class _WireReader: section = self.message.sections[section_number] for _ in range(qcount): qname = self.parser.get_name(self.message.origin) - (rdtype, rdclass) = self.parser.get_struct('!HH') - (rdclass, rdtype, _, _) = \ - self.message._parse_rr_header(section_number, qname, rdclass, - rdtype) - self.message.find_rrset(section, qname, rdclass, rdtype, - create=True, force_unique=True) + (rdtype, rdclass) = self.parser.get_struct("!HH") + (rdclass, rdtype, _, _) = self.message._parse_rr_header( + section_number, qname, rdclass, rdtype + ) + self.message.find_rrset( + section, qname, rdclass, rdtype, create=True, force_unique=True + ) def _add_error(self, e): self.errors.append(MessageError(e, self.parser.current)) @@ -964,16 +1040,20 @@ class _WireReader: name = absolute_name.relativize(self.message.origin) else: name = absolute_name - (rdtype, rdclass, ttl, rdlen) = self.parser.get_struct('!HHIH') + (rdtype, rdclass, ttl, rdlen) = self.parser.get_struct("!HHIH") if rdtype in (dns.rdatatype.OPT, dns.rdatatype.TSIG): - (rdclass, rdtype, deleting, empty) = \ - self.message._parse_special_rr_header(section_number, - count, i, name, - rdclass, rdtype) + ( + rdclass, + rdtype, + deleting, + empty, + ) = self.message._parse_special_rr_header( + section_number, count, i, name, rdclass, rdtype + ) else: - (rdclass, rdtype, deleting, empty) = \ - self.message._parse_rr_header(section_number, - name, rdclass, rdtype) + (rdclass, rdtype, deleting, empty) = self.message._parse_rr_header( + section_number, name, rdclass, rdtype + ) try: rdata_start = self.parser.current if empty: @@ -983,9 +1063,9 @@ class _WireReader: covers = dns.rdatatype.NONE else: with self.parser.restrict_to(rdlen): - rd = dns.rdata.from_wire_parser(rdclass, rdtype, - self.parser, - self.message.origin) + rd = dns.rdata.from_wire_parser( + rdclass, rdtype, self.parser, self.message.origin + ) covers = rd.covers() if self.message.xfr and rdtype == dns.rdatatype.SOA: force_unique = True @@ -993,8 +1073,7 @@ class _WireReader: self.message.opt = dns.rrset.from_rdata(name, ttl, rd) elif rdtype == dns.rdatatype.TSIG: if self.keyring is None: - raise UnknownTSIGKey('got signed message without ' - 'keyring') + raise UnknownTSIGKey("got signed message without " "keyring") if isinstance(self.keyring, dict): key = self.keyring.get(absolute_name) if isinstance(key, bytes): @@ -1006,25 +1085,31 @@ class _WireReader: if key is None: raise UnknownTSIGKey("key '%s' unknown" % name) self.message.keyring = key - self.message.tsig_ctx = \ - dns.tsig.validate(self.parser.wire, - key, - absolute_name, - rd, - int(time.time()), - self.message.request_mac, - rr_start, - self.message.tsig_ctx, - self.multi) - self.message.tsig = dns.rrset.from_rdata(absolute_name, 0, - rd) + self.message.tsig_ctx = dns.tsig.validate( + self.parser.wire, + key, + absolute_name, + rd, + int(time.time()), + self.message.request_mac, + rr_start, + self.message.tsig_ctx, + self.multi, + ) + self.message.tsig = dns.rrset.from_rdata(absolute_name, 0, rd) else: - rrset = self.message.find_rrset(section, name, - rdclass, rdtype, covers, - deleting, True, - force_unique) + rrset = self.message.find_rrset( + section, + name, + rdclass, + rdtype, + covers, + deleting, + True, + force_unique, + ) if rd is not None: - if ttl > 0x7fffffff: + if ttl > 0x7FFFFFFF: ttl = 0 rrset.add(rd, ttl) except Exception as e: @@ -1040,14 +1125,16 @@ class _WireReader: if self.parser.remaining() < 12: raise ShortHeader - (id, flags, qcount, ancount, aucount, adcount) = \ - self.parser.get_struct('!HHHHHH') + (id, flags, qcount, ancount, aucount, adcount) = self.parser.get_struct( + "!HHHHHH" + ) factory = _message_factory_from_opcode(dns.opcode.from_flags(flags)) self.message = factory(id=id) self.message.flags = dns.flags.Flag(flags) self.initialize_message(self.message) - self.one_rr_per_rrset = \ - self.message._get_one_rr_per_rrset(self.one_rr_per_rrset) + self.one_rr_per_rrset = self.message._get_one_rr_per_rrset( + self.one_rr_per_rrset + ) try: self._get_question(MessageSection.QUESTION, qcount) if self.question_only: @@ -1057,8 +1144,7 @@ class _WireReader: self._get_section(MessageSection.ADDITIONAL, adcount) if not self.ignore_trailing and self.parser.remaining() != 0: raise TrailingJunk - if self.multi and self.message.tsig_ctx and \ - not self.message.had_tsig: + if self.multi and self.message.tsig_ctx and not self.message.had_tsig: self.message.tsig_ctx.update(self.parser.wire) except Exception as e: if self.continue_on_error: @@ -1068,73 +1154,78 @@ class _WireReader: return self.message -def from_wire(wire: bytes, keyring: Optional[Any]=None, request_mac: Optional[bytes]=b'', - xfr: bool=False, origin: Optional[dns.name.Name]=None, - tsig_ctx: Optional[Union[dns.tsig.HMACTSig, dns.tsig.GSSTSig]]=None, - multi: bool=False, question_only: bool=False, one_rr_per_rrset: bool=False, - ignore_trailing: bool=False, raise_on_truncation: bool=False, - continue_on_error: bool=False) -> Message: +def from_wire( + wire: bytes, + keyring: Optional[Any] = None, + request_mac: Optional[bytes] = b"", + xfr: bool = False, + origin: Optional[dns.name.Name] = None, + tsig_ctx: Optional[Union[dns.tsig.HMACTSig, dns.tsig.GSSTSig]] = None, + multi: bool = False, + question_only: bool = False, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + raise_on_truncation: bool = False, + continue_on_error: bool = False, +) -> Message: """Convert a DNS wire format message into a message object. - *keyring*, a ``dns.tsig.Key`` or ``dict``, the key or keyring to use if the - message is signed. + *keyring*, a ``dns.tsig.Key`` or ``dict``, the key or keyring to use if the message + is signed. - *request_mac*, a ``bytes`` or ``None``. If the message is a response to a TSIG-signed - request, *request_mac* should be set to the MAC of that request. + *request_mac*, a ``bytes`` or ``None``. If the message is a response to a + TSIG-signed request, *request_mac* should be set to the MAC of that request. - *xfr*, a ``bool``, should be set to ``True`` if this message is part of a - zone transfer. + *xfr*, a ``bool``, should be set to ``True`` if this message is part of a zone + transfer. *origin*, a ``dns.name.Name`` or ``None``. If the message is part of a zone - transfer, *origin* should be the origin name of the zone. If not ``None``, - names will be relativized to the origin. + transfer, *origin* should be the origin name of the zone. If not ``None``, names + will be relativized to the origin. - *tsig_ctx*, a ``dns.tsig.HMACTSig`` or ``dns.tsig.GSSTSig`` object, the - ongoing TSIG context, used when validating zone transfers. + *tsig_ctx*, a ``dns.tsig.HMACTSig`` or ``dns.tsig.GSSTSig`` object, the ongoing TSIG + context, used when validating zone transfers. - *multi*, a ``bool``, should be set to ``True`` if this message is part of a - multiple message sequence. + *multi*, a ``bool``, should be set to ``True`` if this message is part of a multiple + message sequence. - *question_only*, a ``bool``. If ``True``, read only up to the end of the - question section. + *question_only*, a ``bool``. If ``True``, read only up to the end of the question + section. - *one_rr_per_rrset*, a ``bool``. If ``True``, put each RR into its own - RRset. + *one_rr_per_rrset*, a ``bool``. If ``True``, put each RR into its own RRset. - *ignore_trailing*, a ``bool``. If ``True``, ignore trailing junk at end of - the message. + *ignore_trailing*, a ``bool``. If ``True``, ignore trailing junk at end of the + message. - *raise_on_truncation*, a ``bool``. If ``True``, raise an exception if the - TC bit is set. + *raise_on_truncation*, a ``bool``. If ``True``, raise an exception if the TC bit is + set. - *continue_on_error*, a ``bool``. If ``True``, try to continue parsing even - if errors occur. Erroneous rdata will be ignored. Errors will be - accumulated as a list of MessageError objects in the message's ``errors`` - attribute. This option is recommended only for DNS analysis tools, or for - use in a server as part of an error handling path. The default is - ``False``. + *continue_on_error*, a ``bool``. If ``True``, try to continue parsing even if + errors occur. Erroneous rdata will be ignored. Errors will be accumulated as a + list of MessageError objects in the message's ``errors`` attribute. This option is + recommended only for DNS analysis tools, or for use in a server as part of an error + handling path. The default is ``False``. - Raises ``dns.message.ShortHeader`` if the message is less than 12 octets - long. + Raises ``dns.message.ShortHeader`` if the message is less than 12 octets long. - Raises ``dns.message.TrailingJunk`` if there were octets in the message past - the end of the proper DNS message, and *ignore_trailing* is ``False``. + Raises ``dns.message.TrailingJunk`` if there were octets in the message past the end + of the proper DNS message, and *ignore_trailing* is ``False``. Raises ``dns.message.BadEDNS`` if an OPT record was in the wrong section, or occurred more than once. - Raises ``dns.message.BadTSIG`` if a TSIG record was not the last record of - the additional data section. + Raises ``dns.message.BadTSIG`` if a TSIG record was not the last record of the + additional data section. - Raises ``dns.message.Truncated`` if the TC flag is set and - *raise_on_truncation* is ``True``. + Raises ``dns.message.Truncated`` if the TC flag is set and *raise_on_truncation* is + ``True``. Returns a ``dns.message.Message``. """ # We permit None for request_mac solely for backwards compatibility if request_mac is None: - request_mac = b'' + request_mac = b"" def initialize_message(message): message.request_mac = request_mac @@ -1142,14 +1233,24 @@ def from_wire(wire: bytes, keyring: Optional[Any]=None, request_mac: Optional[by message.origin = origin message.tsig_ctx = tsig_ctx - reader = _WireReader(wire, initialize_message, question_only, - one_rr_per_rrset, ignore_trailing, keyring, multi, - continue_on_error) + reader = _WireReader( + wire, + initialize_message, + question_only, + one_rr_per_rrset, + ignore_trailing, + keyring, + multi, + continue_on_error, + ) try: m = reader.read() except dns.exception.FormError: - if reader.message and (reader.message.flags & dns.flags.TC) and \ - raise_on_truncation: + if ( + reader.message + and (reader.message.flags & dns.flags.TC) + and raise_on_truncation + ): raise Truncated(message=reader.message) else: raise @@ -1177,8 +1278,15 @@ class _TextReader: relativize_to: the origin to relativize to. """ - def __init__(self, text, idna_codec, one_rr_per_rrset=False, - origin=None, relativize=True, relativize_to=None): + def __init__( + self, + text, + idna_codec, + one_rr_per_rrset=False, + origin=None, + relativize=True, + relativize_to=None, + ): self.message = None self.tok = dns.tokenizer.Tokenizer(text, idna_codec=idna_codec) self.last_name = None @@ -1199,19 +1307,19 @@ class _TextReader: token = self.tok.get() what = token.value - if what == 'id': + if what == "id": self.id = self.tok.get_int() - elif what == 'flags': + elif what == "flags": while True: token = self.tok.get() if not token.is_identifier(): self.tok.unget(token) break self.flags = self.flags | dns.flags.from_text(token.value) - elif what == 'edns': + elif what == "edns": self.edns = self.tok.get_int() self.ednsflags = self.ednsflags | (self.edns << 16) - elif what == 'eflags': + elif what == "eflags": if self.edns < 0: self.edns = 0 while True: @@ -1219,17 +1327,16 @@ class _TextReader: if not token.is_identifier(): self.tok.unget(token) break - self.ednsflags = self.ednsflags | \ - dns.flags.edns_from_text(token.value) - elif what == 'payload': + self.ednsflags = self.ednsflags | dns.flags.edns_from_text(token.value) + elif what == "payload": self.payload = self.tok.get_int() if self.edns < 0: self.edns = 0 - elif what == 'opcode': + elif what == "opcode": text = self.tok.get_string() self.opcode = dns.opcode.from_text(text) self.flags = self.flags | dns.opcode.to_flags(self.opcode) - elif what == 'rcode': + elif what == "rcode": text = self.tok.get_string() self.rcode = dns.rcode.from_text(text) else: @@ -1242,9 +1349,9 @@ class _TextReader: section = self.message.sections[section_number] token = self.tok.get(want_leading=True) if not token.is_whitespace(): - self.last_name = self.tok.as_name(token, self.message.origin, - self.relativize, - self.relativize_to) + self.last_name = self.tok.as_name( + token, self.message.origin, self.relativize, self.relativize_to + ) name = self.last_name if name is None: raise NoPreviousName @@ -1263,10 +1370,12 @@ class _TextReader: rdclass = dns.rdataclass.IN # Type rdtype = dns.rdatatype.from_text(token.value) - (rdclass, rdtype, _, _) = \ - self.message._parse_rr_header(section_number, name, rdclass, rdtype) - self.message.find_rrset(section, name, rdclass, rdtype, create=True, - force_unique=True) + (rdclass, rdtype, _, _) = self.message._parse_rr_header( + section_number, name, rdclass, rdtype + ) + self.message.find_rrset( + section, name, rdclass, rdtype, create=True, force_unique=True + ) self.tok.get_eol() def _rr_line(self, section_number): @@ -1278,9 +1387,9 @@ class _TextReader: # Name token = self.tok.get(want_leading=True) if not token.is_whitespace(): - self.last_name = self.tok.as_name(token, self.message.origin, - self.relativize, - self.relativize_to) + self.last_name = self.tok.as_name( + token, self.message.origin, self.relativize, self.relativize_to + ) name = self.last_name if name is None: raise NoPreviousName @@ -1309,8 +1418,9 @@ class _TextReader: rdclass = dns.rdataclass.IN # Type rdtype = dns.rdatatype.from_text(token.value) - (rdclass, rdtype, deleting, empty) = \ - self.message._parse_rr_header(section_number, name, rdclass, rdtype) + (rdclass, rdtype, deleting, empty) = self.message._parse_rr_header( + section_number, name, rdclass, rdtype + ) token = self.tok.get() if empty and not token.is_eol_or_eof(): raise dns.exception.SyntaxError @@ -1318,16 +1428,28 @@ class _TextReader: raise dns.exception.UnexpectedEnd if not token.is_eol_or_eof(): self.tok.unget(token) - rd = dns.rdata.from_text(rdclass, rdtype, self.tok, - self.message.origin, self.relativize, - self.relativize_to) + rd = dns.rdata.from_text( + rdclass, + rdtype, + self.tok, + self.message.origin, + self.relativize, + self.relativize_to, + ) covers = rd.covers() else: rd = None covers = dns.rdatatype.NONE - rrset = self.message.find_rrset(section, name, - rdclass, rdtype, covers, - deleting, True, self.one_rr_per_rrset) + rrset = self.message.find_rrset( + section, + name, + rdclass, + rdtype, + covers, + deleting, + True, + self.one_rr_per_rrset, + ) if rd is not None: rrset.add(rd, ttl) @@ -1355,7 +1477,7 @@ class _TextReader: break if token.is_comment(): u = token.value.upper() - if u == 'HEADER': + if u == "HEADER": line_method = self._header_line if self.message: @@ -1370,8 +1492,9 @@ class _TextReader: # use the one we just created. if not self.message: self.message = message - self.one_rr_per_rrset = \ - message._get_one_rr_per_rrset(self.one_rr_per_rrset) + self.one_rr_per_rrset = message._get_one_rr_per_rrset( + self.one_rr_per_rrset + ) if section_number == MessageSection.QUESTION: line_method = self._question_line else: @@ -1388,9 +1511,14 @@ class _TextReader: return self.message -def from_text(text: str, idna_codec: Optional[dns.name.IDNACodec]=None, - one_rr_per_rrset: bool=False, origin: Optional[dns.name.Name]=None, - relativize: bool=True, relativize_to: Optional[dns.name.Name]=None) -> Message: +def from_text( + text: str, + idna_codec: Optional[dns.name.IDNACodec] = None, + one_rr_per_rrset: bool = False, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, +) -> Message: """Convert the text format message into a message object. The reader stops after reading the first blank line in the input to @@ -1425,12 +1553,17 @@ def from_text(text: str, idna_codec: Optional[dns.name.IDNACodec]=None, # since it's an implementation detail. The official file # interface is from_file(). - reader = _TextReader(text, idna_codec, one_rr_per_rrset, origin, - relativize, relativize_to) + reader = _TextReader( + text, idna_codec, one_rr_per_rrset, origin, relativize, relativize_to + ) return reader.read() -def from_file(f: Any, idna_codec: Optional[dns.name.IDNACodec]=None, one_rr_per_rrset: bool=False) -> Message: +def from_file( + f: Any, + idna_codec: Optional[dns.name.IDNACodec] = None, + one_rr_per_rrset: bool = False, +) -> Message: """Read the next text format message from the specified file. Message blocks are separated by a single blank line. @@ -1459,14 +1592,20 @@ def from_file(f: Any, idna_codec: Optional[dns.name.IDNACodec]=None, one_rr_per_ assert False # for mypy lgtm[py/unreachable-statement] -def make_query(qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - use_edns: Optional[Union[int, bool]]=None, - want_dnssec: bool=False, ednsflags: Optional[int]=None, payload: Optional[int]=None, - request_payload: Optional[int]=None, options: Optional[List[dns.edns.Option]]=None, - idna_codec: Optional[dns.name.IDNACodec]=None, id: Optional[int]=None, - flags: int=dns.flags.RD) -> QueryMessage: +def make_query( + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + use_edns: Optional[Union[int, bool]] = None, + want_dnssec: bool = False, + ednsflags: Optional[int] = None, + payload: Optional[int] = None, + request_payload: Optional[int] = None, + options: Optional[List[dns.edns.Option]] = None, + idna_codec: Optional[dns.name.IDNACodec] = None, + id: Optional[int] = None, + flags: int = dns.flags.RD, +) -> QueryMessage: """Make a query message. The query name, type, and class may all be specified either @@ -1523,30 +1662,36 @@ def make_query(qname: Union[dns.name.Name, str], the_rdclass = dns.rdataclass.RdataClass.make(rdclass) m = QueryMessage(id=id) m.flags = dns.flags.Flag(flags) - m.find_rrset(m.question, qname, the_rdclass, the_rdtype, create=True, - force_unique=True) + m.find_rrset( + m.question, qname, the_rdclass, the_rdtype, create=True, force_unique=True + ) # only pass keywords on to use_edns if they have been set to a # non-None value. Setting a field will turn EDNS on if it hasn't # been configured. kwargs: Dict[str, Any] = {} if ednsflags is not None: - kwargs['ednsflags'] = ednsflags + kwargs["ednsflags"] = ednsflags if payload is not None: - kwargs['payload'] = payload + kwargs["payload"] = payload if request_payload is not None: - kwargs['request_payload'] = request_payload + kwargs["request_payload"] = request_payload if options is not None: - kwargs['options'] = options + kwargs["options"] = options if kwargs and use_edns is None: use_edns = 0 - kwargs['edns'] = use_edns + kwargs["edns"] = use_edns m.use_edns(**kwargs) m.want_dnssec(want_dnssec) return m -def make_response(query: Message, recursion_available: bool=False, our_payload: int=8192, - fudge: int=300, tsig_error: int=0) -> Message: +def make_response( + query: Message, + recursion_available: bool = False, + our_payload: int = 8192, + fudge: int = 300, + tsig_error: int = 0, +) -> Message: """Make a message which is a response for the specified query. The message returned is really a response skeleton; it has all of the infrastructure required of a response, but none of the @@ -1573,7 +1718,7 @@ def make_response(query: Message, recursion_available: bool=False, our_payload: """ if query.flags & dns.flags.QR: - raise dns.exception.FormError('specified query message is not a query') + raise dns.exception.FormError("specified query message is not a query") factory = _message_factory_from_opcode(query.opcode()) response = factory(id=query.id) response.flags = dns.flags.QR | (query.flags & dns.flags.RD) @@ -1584,11 +1729,19 @@ def make_response(query: Message, recursion_available: bool=False, our_payload: if query.edns >= 0: response.use_edns(0, 0, our_payload, query.payload) if query.had_tsig: - response.use_tsig(query.keyring, query.keyname, fudge, None, - tsig_error, b'', query.keyalgorithm) + response.use_tsig( + query.keyring, + query.keyname, + fudge, + None, + tsig_error, + b"", + query.keyalgorithm, + ) response.request_mac = query.mac return response + ### BEGIN generated MessageSection constants QUESTION = MessageSection.QUESTION diff --git a/dns/name.py b/dns/name.py index daf1259..2ebda4a 100644 --- a/dns/name.py +++ b/dns/name.py @@ -23,9 +23,11 @@ from typing import Any, Dict, Iterable, Optional, Tuple, Union import copy import struct -import encodings.idna # type: ignore +import encodings.idna # type: ignore + try: - import idna # type: ignore + import idna # type: ignore + have_idna_2008 = True except ImportError: # pragma: no cover have_idna_2008 = False @@ -36,7 +38,7 @@ import dns.exception import dns.immutable -CompressType = Dict['Name', int] +CompressType = Dict["Name", int] class NameRelation(dns.enum.IntEnum): @@ -111,6 +113,7 @@ class NoParent(dns.exception.DNSException): """An attempt was made to get the parent of the root name or the empty name.""" + class NoIDNA2008(dns.exception.DNSException): """IDNA 2008 processing was requested but the idna module is not available.""" @@ -119,10 +122,11 @@ class NoIDNA2008(dns.exception.DNSException): class IDNAException(dns.exception.DNSException): """IDNA processing raised an exception.""" - supp_kwargs = {'idna_exception'} + supp_kwargs = {"idna_exception"} fmt = "IDNA processing exception: {idna_exception}" - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -130,6 +134,7 @@ class IDNAException(dns.exception.DNSException): _escaped = b'"().;\\@$' _escaped_text = '"().;\\@$' + def _escapify(label: Union[bytes, str]) -> str: """Escape the characters in label which need it. @returns: the escaped string @@ -137,23 +142,23 @@ def _escapify(label: Union[bytes, str]) -> str: if isinstance(label, bytes): # Ordinary DNS label mode. Escape special characters and values # < 0x20 or > 0x7f. - text = '' + text = "" for c in label: if c in _escaped: - text += '\\' + chr(c) + text += "\\" + chr(c) elif c > 0x20 and c < 0x7F: text += chr(c) else: - text += '\\%03d' % c + text += "\\%03d" % c return text # Unicode label mode. Escape only special characters and values < 0x20 - text = '' + text = "" for uc in label: if uc in _escaped_text: - text += '\\' + uc - elif uc <= '\x20': - text += '\\%03d' % ord(uc) + text += "\\" + uc + elif uc <= "\x20": + text += "\\%03d" % ord(uc) else: text += uc return text @@ -166,7 +171,7 @@ class IDNACodec: pass def is_idna(self, label: bytes) -> bool: - return label.lower().startswith(b'xn--') + return label.lower().startswith(b"xn--") def encode(self, label: str) -> bytes: raise NotImplementedError # pragma: no cover @@ -175,7 +180,7 @@ class IDNACodec: # We do not apply any IDNA policy on decode. if self.is_idna(label): try: - slabel = label[4:].decode('punycode') + slabel = label[4:].decode("punycode") return _escapify(slabel) except Exception as e: raise IDNAException(idna_exception=e) @@ -186,7 +191,7 @@ class IDNACodec: class IDNA2003Codec(IDNACodec): """IDNA 2003 encoder/decoder.""" - def __init__(self, strict_decode: bool=False): + def __init__(self, strict_decode: bool = False): """Initialize the IDNA 2003 encoder/decoder. *strict_decode* is a ``bool``. If `True`, then IDNA2003 checking @@ -200,8 +205,8 @@ class IDNA2003Codec(IDNACodec): def encode(self, label: str) -> bytes: """Encode *label*.""" - if label == '': - return b'' + if label == "": + return b"" try: return encodings.idna.ToASCII(label) except UnicodeError: @@ -211,8 +216,8 @@ class IDNA2003Codec(IDNACodec): """Decode *label*.""" if not self.strict_decode: return super().decode(label) - if label == b'': - return '' + if label == b"": + return "" try: return _escapify(encodings.idna.ToUnicode(label)) except Exception as e: @@ -220,11 +225,15 @@ class IDNA2003Codec(IDNACodec): class IDNA2008Codec(IDNACodec): - """IDNA 2008 encoder/decoder. - """ - - def __init__(self, uts_46: bool=False, transitional: bool=False, - allow_pure_ascii: bool=False, strict_decode: bool=False): + """IDNA 2008 encoder/decoder.""" + + def __init__( + self, + uts_46: bool = False, + transitional: bool = False, + allow_pure_ascii: bool = False, + strict_decode: bool = False, + ): """Initialize the IDNA 2008 encoder/decoder. *uts_46* is a ``bool``. If True, apply Unicode IDNA @@ -254,10 +263,10 @@ class IDNA2008Codec(IDNACodec): self.strict_decode = strict_decode def encode(self, label: str) -> bytes: - if label == '': - return b'' + if label == "": + return b"" if self.allow_pure_ascii and is_all_ascii(label): - encoded = label.encode('ascii') + encoded = label.encode("ascii") if len(encoded) > 63: raise LabelTooLong return encoded @@ -268,7 +277,7 @@ class IDNA2008Codec(IDNACodec): label = idna.uts46_remap(label, False, self.transitional) return idna.alabel(label) except idna.IDNAError as e: - if e.args[0] == 'Label too long': + if e.args[0] == "Label too long": raise LabelTooLong else: raise IDNAException(idna_exception=e) @@ -276,8 +285,8 @@ class IDNA2008Codec(IDNACodec): def decode(self, label: bytes) -> str: if not self.strict_decode: return super().decode(label) - if label == b'': - return '' + if label == b"": + return "" if not have_idna_2008: raise NoIDNA2008 try: @@ -288,6 +297,7 @@ class IDNA2008Codec(IDNACodec): except (idna.IDNAError, UnicodeError) as e: raise IDNAException(idna_exception=e) + IDNA_2003_Practical = IDNA2003Codec(False) IDNA_2003_Strict = IDNA2003Codec(True) IDNA_2003 = IDNA_2003_Practical @@ -297,6 +307,7 @@ IDNA_2008_Strict = IDNA2008Codec(False, False, False, True) IDNA_2008_Transitional = IDNA2008Codec(True, True, False, False) IDNA_2008 = IDNA_2008_Practical + def _validate_labels(labels: Tuple[bytes, ...]) -> None: """Check for empty labels in the middle of a label sequence, labels that are too long, and for too many labels. @@ -318,7 +329,7 @@ def _validate_labels(labels: Tuple[bytes, ...]) -> None: total += ll + 1 if ll > 63: raise LabelTooLong - if i < 0 and label == b'': + if i < 0 and label == b"": i = j j += 1 if total > 255: @@ -350,11 +361,10 @@ class Name: of the class are immutable. """ - __slots__ = ['labels'] + __slots__ = ["labels"] def __init__(self, labels: Iterable[Union[bytes, str]]): - """*labels* is any iterable whose values are ``str`` or ``bytes``. - """ + """*labels* is any iterable whose values are ``str`` or ``bytes``.""" blabels = [_maybe_convert_to_binary(x) for x in labels] self.labels = tuple(blabels) @@ -368,10 +378,10 @@ class Name: def __getstate__(self): # Names can be pickled - return {'labels': self.labels} + return {"labels": self.labels} def __setstate__(self, state): - super().__setattr__('labels', state['labels']) + super().__setattr__("labels", state["labels"]) _validate_labels(self.labels) def is_absolute(self) -> bool: @@ -380,7 +390,7 @@ class Name: Returns a ``bool``. """ - return len(self.labels) > 0 and self.labels[-1] == b'' + return len(self.labels) > 0 and self.labels[-1] == b"" def is_wild(self) -> bool: """Is this name wild? (I.e. Is the least significant label '*'?) @@ -388,7 +398,7 @@ class Name: Returns a ``bool``. """ - return len(self.labels) > 0 and self.labels[0] == b'*' + return len(self.labels) > 0 and self.labels[0] == b"*" def __hash__(self) -> int: """Return a case-insensitive hash of the name. @@ -402,7 +412,7 @@ class Name: h += (h << 3) + c return h - def fullcompare(self, other: 'Name') -> Tuple[NameRelation, int, int]: + def fullcompare(self, other: "Name") -> Tuple[NameRelation, int, int]: """Compare two names, returning a 3-tuple ``(relation, order, nlabels)``. @@ -478,7 +488,7 @@ class Name: namereln = NameRelation.EQUAL return (namereln, order, nlabels) - def is_subdomain(self, other: 'Name') -> bool: + def is_subdomain(self, other: "Name") -> bool: """Is self a subdomain of other? Note that the notion of subdomain includes equality, e.g. @@ -492,7 +502,7 @@ class Name: return True return False - def is_superdomain(self, other: 'Name') -> bool: + def is_superdomain(self, other: "Name") -> bool: """Is self a superdomain of other? Note that the notion of superdomain includes equality, e.g. @@ -506,7 +516,7 @@ class Name: return True return False - def canonicalize(self) -> 'Name': + def canonicalize(self) -> "Name": """Return a name which is equal to the current name, but is in DNSSEC canonical form. """ @@ -550,12 +560,12 @@ class Name: return NotImplemented def __repr__(self): - return '<DNS name ' + self.__str__() + '>' + return "<DNS name " + self.__str__() + ">" def __str__(self): return self.to_text(False) - def to_text(self, omit_final_dot: bool=False) -> str: + def to_text(self, omit_final_dot: bool = False) -> str: """Convert name to DNS text format. *omit_final_dot* is a ``bool``. If True, don't emit the final @@ -566,17 +576,19 @@ class Name: """ if len(self.labels) == 0: - return '@' - if len(self.labels) == 1 and self.labels[0] == b'': - return '.' + return "@" + if len(self.labels) == 1 and self.labels[0] == b"": + return "." if omit_final_dot and self.is_absolute(): l = self.labels[:-1] else: l = self.labels - s = '.'.join(map(_escapify, l)) + s = ".".join(map(_escapify, l)) return s - def to_unicode(self, omit_final_dot: bool=False, idna_codec: Optional[IDNACodec]=None) -> str: + def to_unicode( + self, omit_final_dot: bool = False, idna_codec: Optional[IDNACodec] = None + ) -> str: """Convert name to Unicode text format. IDN ACE labels are converted to Unicode. @@ -595,18 +607,18 @@ class Name: """ if len(self.labels) == 0: - return '@' - if len(self.labels) == 1 and self.labels[0] == b'': - return '.' + return "@" + if len(self.labels) == 1 and self.labels[0] == b"": + return "." if omit_final_dot and self.is_absolute(): l = self.labels[:-1] else: l = self.labels if idna_codec is None: idna_codec = IDNA_2003_Practical - return '.'.join([idna_codec.decode(x) for x in l]) + return ".".join([idna_codec.decode(x) for x in l]) - def to_digestable(self, origin: Optional['Name']=None) -> bytes: + def to_digestable(self, origin: Optional["Name"] = None) -> bytes: """Convert name to a format suitable for digesting in hashes. The name is canonicalized and converted to uncompressed wire @@ -627,8 +639,13 @@ class Name: assert digest is not None return digest - def to_wire(self, file: Optional[Any]=None, compress: Optional[CompressType]=None, - origin: Optional['Name']=None, canonicalize: bool=False) -> Optional[bytes]: + def to_wire( + self, + file: Optional[Any] = None, + compress: Optional[CompressType] = None, + origin: Optional["Name"] = None, + canonicalize: bool = False, + ) -> Optional[bytes]: """Convert name to wire format, possibly compressing it. *file* is the file where the name is emitted (typically an @@ -691,17 +708,17 @@ class Name: else: pos = None if pos is not None: - value = 0xc000 + pos - s = struct.pack('!H', value) + value = 0xC000 + pos + s = struct.pack("!H", value) file.write(s) break else: if compress is not None and len(n) > 1: pos = file.tell() - if pos <= 0x3fff: + if pos <= 0x3FFF: compress[n] = pos l = len(label) - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) if l > 0: if canonicalize: file.write(label.lower()) @@ -726,7 +743,7 @@ class Name: def __sub__(self, other): return self.relativize(other) - def split(self, depth: int) -> Tuple['Name', 'Name']: + def split(self, depth: int) -> Tuple["Name", "Name"]: """Split a name into a prefix and suffix names at the specified depth. *depth* is an ``int`` specifying the number of labels in the suffix @@ -743,11 +760,10 @@ class Name: elif depth == l: return (dns.name.empty, self) elif depth < 0 or depth > l: - raise ValueError( - 'depth must be >= 0 and <= the length of the name') - return (Name(self[: -depth]), Name(self[-depth:])) + raise ValueError("depth must be >= 0 and <= the length of the name") + return (Name(self[:-depth]), Name(self[-depth:])) - def concatenate(self, other: 'Name') -> 'Name': + def concatenate(self, other: "Name") -> "Name": """Return a new name which is the concatenation of self and other. Raises ``dns.name.AbsoluteConcatenation`` if the name is @@ -762,7 +778,7 @@ class Name: labels.extend(list(other.labels)) return Name(labels) - def relativize(self, origin: 'Name') -> 'Name': + def relativize(self, origin: "Name") -> "Name": """If the name is a subdomain of *origin*, return a new name which is the name relative to origin. Otherwise return the name. @@ -778,7 +794,7 @@ class Name: else: return self - def derelativize(self, origin: 'Name') -> 'Name': + def derelativize(self, origin: "Name") -> "Name": """If the name is a relative name, return a new name which is the concatenation of the name and origin. Otherwise return the name. @@ -794,7 +810,9 @@ class Name: else: return self - def choose_relativity(self, origin: Optional['Name']=None, relativize: bool=True) -> 'Name': + def choose_relativity( + self, origin: Optional["Name"] = None, relativize: bool = True + ) -> "Name": """Return a name with the relativity desired by the caller. If *origin* is ``None``, then the name is returned. @@ -813,7 +831,7 @@ class Name: else: return self - def parent(self) -> 'Name': + def parent(self) -> "Name": """Return the parent of the name. For example, the parent of ``www.dnspython.org.`` is ``dnspython.org``. @@ -828,13 +846,17 @@ class Name: raise NoParent return Name(self.labels[1:]) + #: The root name, '.' -root = Name([b'']) +root = Name([b""]) #: The empty name. empty = Name([]) -def from_unicode(text: str, origin: Optional[Name]=root, idna_codec: Optional[IDNACodec]=None) -> Name: + +def from_unicode( + text: str, origin: Optional[Name] = root, idna_codec: Optional[IDNACodec] = None +) -> Name: """Convert unicode text into a Name object. Labels are encoded in IDN ACE form according to rules specified by @@ -857,17 +879,17 @@ def from_unicode(text: str, origin: Optional[Name]=root, idna_codec: Optional[ID if not (origin is None or isinstance(origin, Name)): raise ValueError("origin must be a Name or None") labels = [] - label = '' + label = "" escaping = False edigits = 0 total = 0 if idna_codec is None: idna_codec = IDNA_2003 - if text == '@': - text = '' + if text == "@": + text = "" if text: - if text in ['.', '\u3002', '\uff0e', '\uff61']: - return Name([b'']) # no Unicode "u" on this constant! + if text in [".", "\u3002", "\uff0e", "\uff61"]: + return Name([b""]) # no Unicode "u" on this constant! for c in text: if escaping: if edigits == 0: @@ -886,12 +908,12 @@ def from_unicode(text: str, origin: Optional[Name]=root, idna_codec: Optional[ID if edigits == 3: escaping = False label += chr(total) - elif c in ['.', '\u3002', '\uff0e', '\uff61']: + elif c in [".", "\u3002", "\uff0e", "\uff61"]: if len(label) == 0: raise EmptyLabel labels.append(idna_codec.encode(label)) - label = '' - elif c == '\\': + label = "" + elif c == "\\": escaping = True edigits = 0 total = 0 @@ -902,19 +924,25 @@ def from_unicode(text: str, origin: Optional[Name]=root, idna_codec: Optional[ID if len(label) > 0: labels.append(idna_codec.encode(label)) else: - labels.append(b'') + labels.append(b"") - if (len(labels) == 0 or labels[-1] != b'') and origin is not None: + if (len(labels) == 0 or labels[-1] != b"") and origin is not None: labels.extend(list(origin.labels)) return Name(labels) + def is_all_ascii(text: str) -> bool: for c in text: - if ord(c) > 0x7f: + if ord(c) > 0x7F: return False return True -def from_text(text: Union[bytes, str], origin: Optional[Name]=root, idna_codec: Optional[IDNACodec]=None) -> Name: + +def from_text( + text: Union[bytes, str], + origin: Optional[Name] = root, + idna_codec: Optional[IDNACodec] = None, +) -> Name: """Convert text into a Name object. *text*, a ``bytes`` or ``str``, is the text to convert into a name. @@ -941,23 +969,23 @@ def from_text(text: Union[bytes, str], origin: Optional[Name]=root, idna_codec: # # then it's still "all ASCII" even though the domain name has # codepoints > 127. - text = text.encode('ascii') + text = text.encode("ascii") if not isinstance(text, bytes): raise ValueError("input to from_text() must be a string") if not (origin is None or isinstance(origin, Name)): raise ValueError("origin must be a Name or None") labels = [] - label = b'' + label = b"" escaping = False edigits = 0 total = 0 - if text == b'@': - text = b'' + if text == b"@": + text = b"" if text: - if text == b'.': - return Name([b'']) + if text == b".": + return Name([b""]) for c in text: - byte_ = struct.pack('!B', c) + byte_ = struct.pack("!B", c) if escaping: if edigits == 0: if byte_.isdigit(): @@ -974,13 +1002,13 @@ def from_text(text: Union[bytes, str], origin: Optional[Name]=root, idna_codec: edigits += 1 if edigits == 3: escaping = False - label += struct.pack('!B', total) - elif byte_ == b'.': + label += struct.pack("!B", total) + elif byte_ == b".": if len(label) == 0: raise EmptyLabel labels.append(label) - label = b'' - elif byte_ == b'\\': + label = b"" + elif byte_ == b"\\": escaping = True edigits = 0 total = 0 @@ -991,14 +1019,16 @@ def from_text(text: Union[bytes, str], origin: Optional[Name]=root, idna_codec: if len(label) > 0: labels.append(label) else: - labels.append(b'') - if (len(labels) == 0 or labels[-1] != b'') and origin is not None: + labels.append(b"") + if (len(labels) == 0 or labels[-1] != b"") and origin is not None: labels.extend(list(origin.labels)) return Name(labels) + # we need 'dns.wire.Parser' quoted as dns.name and dns.wire depend on each other. -def from_wire_parser(parser: 'dns.wire.Parser') -> Name: + +def from_wire_parser(parser: "dns.wire.Parser") -> Name: """Convert possibly compressed wire format into a Name. *parser* is a dns.wire.Parser. @@ -1019,7 +1049,7 @@ def from_wire_parser(parser: 'dns.wire.Parser') -> Name: if count < 64: labels.append(parser.get_bytes(count)) elif count >= 192: - current = (count & 0x3f) * 256 + parser.get_uint8() + current = (count & 0x3F) * 256 + parser.get_uint8() if current >= biggest_pointer: raise BadPointer biggest_pointer = current @@ -1027,7 +1057,7 @@ def from_wire_parser(parser: 'dns.wire.Parser') -> Name: else: raise BadLabelType count = parser.get_uint8() - labels.append(b'') + labels.append(b"") return Name(labels) diff --git a/dns/namedict.py b/dns/namedict.py index ec0750c..fe118a3 100644 --- a/dns/namedict.py +++ b/dns/namedict.py @@ -62,7 +62,7 @@ class NameDict(MutableMapping): def __setitem__(self, key, value): if not isinstance(key, dns.name.Name): - raise ValueError('NameDict key must be a name') + raise ValueError("NameDict key must be a name") self.__store[key] = value self.__update_max_depth(key) diff --git a/dns/node.py b/dns/node.py index 5270b53..d870a29 100644 --- a/dns/node.py +++ b/dns/node.py @@ -37,26 +37,28 @@ _cname_types = { # "neutral" types can coexist with a CNAME and thus are not "other data" _neutral_types = { - dns.rdatatype.NSEC, # RFC 4035 section 2.5 + dns.rdatatype.NSEC, # RFC 4035 section 2.5 dns.rdatatype.NSEC3, # This is not likely to happen, but not impossible! - dns.rdatatype.KEY, # RFC 4035 section 2.5, RFC 3007 + dns.rdatatype.KEY, # RFC 4035 section 2.5, RFC 3007 } + def _matches_type_or_its_signature(rdtypes, rdtype, covers): - return rdtype in rdtypes or \ - (rdtype == dns.rdatatype.RRSIG and covers in rdtypes) + return rdtype in rdtypes or (rdtype == dns.rdatatype.RRSIG and covers in rdtypes) @enum.unique class NodeKind(enum.Enum): - """Rdatasets in nodes - """ - REGULAR = 0 # a.k.a "other data" + """Rdatasets in nodes""" + + REGULAR = 0 # a.k.a "other data" NEUTRAL = 1 CNAME = 2 @classmethod - def classify(cls, rdtype: dns.rdatatype.RdataType, covers: dns.rdatatype.RdataType) -> 'NodeKind': + def classify( + cls, rdtype: dns.rdatatype.RdataType, covers: dns.rdatatype.RdataType + ) -> "NodeKind": if _matches_type_or_its_signature(_cname_types, rdtype, covers): return NodeKind.CNAME elif _matches_type_or_its_signature(_neutral_types, rdtype, covers): @@ -65,7 +67,7 @@ class NodeKind(enum.Enum): return NodeKind.REGULAR @classmethod - def classify_rdataset(cls, rdataset: dns.rdataset.Rdataset) -> 'NodeKind': + def classify_rdataset(cls, rdataset: dns.rdataset.Rdataset) -> "NodeKind": return cls.classify(rdataset.rdtype, rdataset.covers) @@ -86,7 +88,7 @@ class Node: deleted. """ - __slots__ = ['rdatasets'] + __slots__ = ["rdatasets"] def __init__(self): # the set of rdatasets, represented as a list. @@ -109,11 +111,11 @@ class Node: for rds in self.rdatasets: if len(rds) > 0: s.write(rds.to_text(name, **kw)) # type: ignore[arg-type] - s.write('\n') + s.write("\n") return s.getvalue()[:-1] def __repr__(self): - return '<DNS node ' + str(id(self)) + '>' + return "<DNS node " + str(id(self)) + ">" def __eq__(self, other): # @@ -149,22 +151,28 @@ class Node: if len(self.rdatasets) > 0: kind = NodeKind.classify_rdataset(rdataset) if kind == NodeKind.CNAME: - self.rdatasets = [rds for rds in self.rdatasets if - NodeKind.classify_rdataset(rds) != - NodeKind.REGULAR] + self.rdatasets = [ + rds + for rds in self.rdatasets + if NodeKind.classify_rdataset(rds) != NodeKind.REGULAR + ] elif kind == NodeKind.REGULAR: - self.rdatasets = [rds for rds in self.rdatasets if - NodeKind.classify_rdataset(rds) != - NodeKind.CNAME] + self.rdatasets = [ + rds + for rds in self.rdatasets + if NodeKind.classify_rdataset(rds) != NodeKind.CNAME + ] # Otherwise the rdataset is NodeKind.NEUTRAL and we do not need to # edit self.rdatasets. self.rdatasets.append(rdataset) - def find_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> dns.rdataset.Rdataset: + def find_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> dns.rdataset.Rdataset: """Find an rdataset matching the specified properties in the current node. @@ -199,11 +207,13 @@ class Node: self._append_rdataset(rds) return rds - def get_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> Optional[dns.rdataset.Rdataset]: """Get an rdataset matching the specified properties in the current node. @@ -234,10 +244,12 @@ class Node: rds = None return rds - def delete_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE) -> None: + def delete_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + ) -> None: """Delete the rdataset matching the specified properties in the current node. @@ -270,13 +282,14 @@ class Node: """ if not isinstance(replacement, dns.rdataset.Rdataset): - raise ValueError('replacement is not an rdataset') + raise ValueError("replacement is not an rdataset") if isinstance(replacement, dns.rrset.RRset): # RRsets are not good replacements as the match() method # is not compatible. replacement = replacement.to_rdataset() - self.delete_rdataset(replacement.rdclass, replacement.rdtype, - replacement.covers) + self.delete_rdataset( + replacement.rdclass, replacement.rdtype, replacement.covers + ) self._append_rdataset(replacement) def classify(self) -> NodeKind: @@ -312,28 +325,34 @@ class ImmutableNode(Node): [dns.rdataset.ImmutableRdataset(rds) for rds in node.rdatasets] ) - def find_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> dns.rdataset.Rdataset: + def find_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> dns.rdataset.Rdataset: if create: raise TypeError("immutable") return super().find_rdataset(rdclass, rdtype, covers, False) - def get_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> Optional[dns.rdataset.Rdataset]: if create: raise TypeError("immutable") return super().get_rdataset(rdclass, rdtype, covers, False) - def delete_rdataset(self, - rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE) -> None: + def delete_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + ) -> None: raise TypeError("immutable") def replace_rdataset(self, replacement: dns.rdataset.Rdataset) -> None: diff --git a/dns/opcode.py b/dns/opcode.py index 971b62c..78b43d2 100644 --- a/dns/opcode.py +++ b/dns/opcode.py @@ -20,6 +20,7 @@ import dns.enum import dns.exception + class Opcode(dns.enum.IntEnum): #: Query QUERY = 0 @@ -104,6 +105,7 @@ def is_update(flags: int) -> bool: return from_flags(flags) == Opcode.UPDATE + ### BEGIN generated Opcode constants QUERY = Opcode.QUERY diff --git a/dns/query.py b/dns/query.py index 09d5107..2c3da4f 100644 --- a/dns/query.py +++ b/dns/query.py @@ -46,6 +46,7 @@ try: import requests from requests_toolbelt.adapters.source import SourceAddressAdapter from requests_toolbelt.adapters.host_header_ssl import HostHeaderSSLAdapter + _have_requests = True except ImportError: # pragma: no cover _have_requests = False @@ -54,6 +55,7 @@ _have_httpx = False _have_http2 = False try: import httpx + _have_httpx = True try: # See if http2 support is available. @@ -69,8 +71,8 @@ have_doh = _have_requests or _have_httpx try: import ssl except ImportError: # pragma: no cover - class ssl: # type: ignore + class ssl: # type: ignore class WantReadException(Exception): pass @@ -85,12 +87,14 @@ except ImportError: # pragma: no cover @classmethod def create_default_context(cls, *args, **kwargs): - raise Exception('no ssl support') + raise Exception("no ssl support") + # Function used to create a socket. Can be overridden if needed in special # situations. socket_factory = socket.socket + class UnexpectedSource(dns.exception.DNSException): """A DNS query response came from an unexpected address or port.""" @@ -151,7 +155,8 @@ def _set_selector_class(selector_class): _selector_class = selector_class -if hasattr(selectors, 'PollSelector'): + +if hasattr(selectors, "PollSelector"): # Prefer poll() on platforms that support it because it has no # limits on the maximum value of a file descriptor (plus it will # be more efficient for high values). @@ -188,18 +193,20 @@ def _matches_destination(af, from_address, destination, ignore_unexpected): # sent to destination. if not destination: return True - if _addresses_equal(af, from_address, destination) or \ - (dns.inet.is_multicast(destination[0]) and - from_address[1:] == destination[1:]): + if _addresses_equal(af, from_address, destination) or ( + dns.inet.is_multicast(destination[0]) and from_address[1:] == destination[1:] + ): return True elif ignore_unexpected: return False - raise UnexpectedSource(f'got a response from {from_address} instead of ' - f'{destination}') + raise UnexpectedSource( + f"got a response from {from_address} instead of " f"{destination}" + ) -def _destination_and_source(where, port, source, source_port, - where_must_be_address=True): +def _destination_and_source( + where, port, source, source_port, where_must_be_address=True +): # Apply defaults and compute destination and source tuples # suitable for use in connect(), sendto(), or bind(). af = None @@ -216,8 +223,9 @@ def _destination_and_source(where, port, source, source_port, if af: # We know the destination af, so source had better agree! if saf != af: - raise ValueError('different address families for source ' + - 'and destination') + raise ValueError( + "different address families for source " + "and destination" + ) else: # We didn't know the destination af, but we know the source, # so that's our af. @@ -227,12 +235,11 @@ def _destination_and_source(where, port, source, source_port, # need to return a source, and we need to use the appropriate # wildcard address as the address. if af == socket.AF_INET: - source = '0.0.0.0' + source = "0.0.0.0" elif af == socket.AF_INET6: - source = '::' + source = "::" else: - raise ValueError('source_port specified but address family is ' - 'unknown') + raise ValueError("source_port specified but address family is " "unknown") # Convert high-level (address, port) tuples into low-level address # tuples. if destination: @@ -241,6 +248,7 @@ def _destination_and_source(where, port, source, source_port, source = dns.inet.low_level_address_tuple((source, source_port), af) return (af, destination, source) + def _make_socket(af, type, source, ssl_context=None, server_hostname=None): s = socket_factory(af, type) try: @@ -249,19 +257,33 @@ def _make_socket(af, type, source, ssl_context=None, server_hostname=None): s.bind(source) if ssl_context: # LGTM gets a false positive here, as our default context is OK - return ssl_context.wrap_socket(s, do_handshake_on_connect=False, # lgtm[py/insecure-protocol] - server_hostname=server_hostname) + return ssl_context.wrap_socket( + s, + do_handshake_on_connect=False, # lgtm[py/insecure-protocol] + server_hostname=server_hostname, + ) else: return s except Exception: s.close() raise -def https(q: dns.message.Message, where: str, timeout: Optional[float]=None, - port: int=443, source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - session: Optional[Any]=None, path: str='/dns-query', post: bool=True, - bootstrap_address: Optional[str]=None, verify: bool=True) -> dns.message.Message: + +def https( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 443, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + session: Optional[Any] = None, + path: str = "/dns-query", + post: bool = True, + bootstrap_address: Optional[str] = None, + verify: bool = True, +) -> dns.message.Message: """Return the response obtained after sending a query via DNS-over-HTTPS. *q*, a ``dns.message.Message``, the query to send. @@ -304,29 +326,26 @@ def https(q: dns.message.Message, where: str, timeout: Optional[float]=None, """ if not have_doh: - raise NoDOH('Neither httpx nor requests is available.') # pragma: no cover + raise NoDOH("Neither httpx nor requests is available.") # pragma: no cover _httpx_ok = _have_httpx wire = q.to_wire() - (af, _, source) = _destination_and_source(where, port, source, source_port, - False) + (af, _, source) = _destination_and_source(where, port, source, source_port, False) transport_adapter = None transport = None - headers = { - "accept": "application/dns-message" - } + headers = {"accept": "application/dns-message"} if af is not None: if af == socket.AF_INET: - url = 'https://{}:{}{}'.format(where, port, path) + url = "https://{}:{}{}".format(where, port, path) elif af == socket.AF_INET6: - url = 'https://[{}]:{}{}'.format(where, port, path) + url = "https://[{}]:{}{}".format(where, port, path) elif bootstrap_address is not None: _httpx_ok = False split_url = urllib.parse.urlsplit(where) if split_url.hostname is None: - raise ValueError('DoH URL has no hostname') - headers['Host'] = split_url.hostname + raise ValueError("DoH URL has no hostname") + headers["Host"] = split_url.hostname url = where.replace(split_url.hostname, bootstrap_address) if _have_requests: transport_adapter = HostHeaderSSLAdapter() @@ -348,22 +367,29 @@ def https(q: dns.message.Message, where: str, timeout: Optional[float]=None, else: _is_httpx = False if _is_httpx and not _httpx_ok: - raise NoDOH('Session is httpx, but httpx cannot be used for ' - 'the requested operation.') + raise NoDOH( + "Session is httpx, but httpx cannot be used for " + "the requested operation." + ) else: _is_httpx = _httpx_ok if not _httpx_ok and not _have_requests: - raise NoDOH('Cannot use httpx for this operation, and ' - 'requests is not available.') + raise NoDOH( + "Cannot use httpx for this operation, and " "requests is not available." + ) with contextlib.ExitStack() as stack: if not session: if _is_httpx: - session = stack.enter_context(httpx.Client(http1=True, - http2=_have_http2, - verify=verify, - transport=transport)) + session = stack.enter_context( + httpx.Client( + http1=True, + http2=_have_http2, + verify=verify, + transport=transport, + ) + ) else: session = stack.enter_context(requests.sessions.Session()) @@ -373,45 +399,56 @@ def https(q: dns.message.Message, where: str, timeout: Optional[float]=None, # see https://tools.ietf.org/html/rfc8484#section-4.1.1 for DoH # GET and POST examples if post: - headers.update({ - "content-type": "application/dns-message", - "content-length": str(len(wire)) - }) + headers.update( + { + "content-type": "application/dns-message", + "content-length": str(len(wire)), + } + ) if _is_httpx: - response = session.post(url, headers=headers, content=wire, - timeout=timeout) + response = session.post( + url, headers=headers, content=wire, timeout=timeout + ) else: - response = session.post(url, headers=headers, data=wire, - timeout=timeout, verify=verify) + response = session.post( + url, headers=headers, data=wire, timeout=timeout, verify=verify + ) else: wire = base64.urlsafe_b64encode(wire).rstrip(b"=") if _is_httpx: twire = wire.decode() # httpx does a repr() if we give it bytes - response = session.get(url, headers=headers, - timeout=timeout, - params={"dns": twire}) + response = session.get( + url, headers=headers, timeout=timeout, params={"dns": twire} + ) else: - response = session.get(url, headers=headers, - timeout=timeout, verify=verify, - params={"dns": wire}) + response = session.get( + url, + headers=headers, + timeout=timeout, + verify=verify, + params={"dns": wire}, + ) # see https://tools.ietf.org/html/rfc8484#section-4.2.1 for info about DoH # status codes if response.status_code < 200 or response.status_code > 299: - raise ValueError('{} responded with status code {}' - '\nResponse body: {}'.format(where, - response.status_code, - response.content)) - r = dns.message.from_wire(response.content, - keyring=q.keyring, - request_mac=q.request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing) + raise ValueError( + "{} responded with status code {}" + "\nResponse body: {}".format(where, response.status_code, response.content) + ) + r = dns.message.from_wire( + response.content, + keyring=q.keyring, + request_mac=q.request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + ) r.time = response.elapsed.total_seconds() if not q.is_response(r): raise BadResponse return r + def _udp_recv(sock, max_size, expiration): """Reads a datagram from the socket. A Timeout exception will be raised if the operation is not completed @@ -439,8 +476,12 @@ def _udp_send(sock, data, destination, expiration): _wait_for_writable(sock, expiration) -def send_udp(sock: Any, what: Union[dns.message.Message, bytes], destination: Any, - expiration: Optional[float]=None) -> Tuple[int, float]: +def send_udp( + sock: Any, + what: Union[dns.message.Message, bytes], + destination: Any, + expiration: Optional[float] = None, +) -> Tuple[int, float]: """Send a DNS message to the specified UDP socket. *sock*, a ``socket``. @@ -464,10 +505,17 @@ def send_udp(sock: Any, what: Union[dns.message.Message, bytes], destination: An return (n, sent_time) -def receive_udp(sock: Any, destination: Optional[Any]=None, expiration: Optional[float]=None, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, - keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]]=None, request_mac: Optional[bytes]=b'', - ignore_trailing: bool=False, raise_on_truncation: bool=False) -> Any: +def receive_udp( + sock: Any, + destination: Optional[Any] = None, + expiration: Optional[float] = None, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]] = None, + request_mac: Optional[bytes] = b"", + ignore_trailing: bool = False, + raise_on_truncation: bool = False, +) -> Any: """Read a DNS message from a UDP socket. *sock*, a ``socket``. @@ -509,26 +557,41 @@ def receive_udp(sock: Any, destination: Optional[Any]=None, expiration: Optional the message arrived from. """ - wire = b'' + wire = b"" while True: (wire, from_address) = _udp_recv(sock, 65535, expiration) - if _matches_destination(sock.family, from_address, destination, - ignore_unexpected): + if _matches_destination( + sock.family, from_address, destination, ignore_unexpected + ): break received_time = time.time() - r = dns.message.from_wire(wire, keyring=keyring, request_mac=request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing, - raise_on_truncation=raise_on_truncation) + r = dns.message.from_wire( + wire, + keyring=keyring, + request_mac=request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + raise_on_truncation=raise_on_truncation, + ) if destination: return (r, received_time) else: return (r, received_time, from_address) -def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - raise_on_truncation: bool=False, sock: Optional[Any]=None) -> dns.message.Message: + +def udp( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + raise_on_truncation: bool = False, + sock: Optional[Any] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via UDP. *q*, a ``dns.message.Message``, the query to send @@ -568,8 +631,9 @@ def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: """ wire = q.to_wire() - (af, destination, source) = _destination_and_source(where, port, - source, source_port) + (af, destination, source) = _destination_and_source( + where, port, source, source_port + ) (begin_time, expiration) = _compute_times(timeout) with contextlib.ExitStack() as stack: if sock: @@ -577,21 +641,39 @@ def udp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: else: s = stack.enter_context(_make_socket(af, socket.SOCK_DGRAM, source)) send_udp(s, wire, destination, expiration) - (r, received_time) = receive_udp(s, destination, expiration, - ignore_unexpected, one_rr_per_rrset, - q.keyring, q.mac, ignore_trailing, - raise_on_truncation) + (r, received_time) = receive_udp( + s, + destination, + expiration, + ignore_unexpected, + one_rr_per_rrset, + q.keyring, + q.mac, + ignore_trailing, + raise_on_truncation, + ) r.time = received_time - begin_time if not q.is_response(r): raise BadResponse return r - assert False # help mypy figure out we can't get here lgtm[py/unreachable-statement] - -def udp_with_fallback(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - ignore_unexpected: bool=False, one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - udp_sock: Optional[Any]=None, - tcp_sock: Optional[Any]=None) -> Tuple[dns.message.Message, bool]: + assert ( + False # help mypy figure out we can't get here lgtm[py/unreachable-statement] + ) + + +def udp_with_fallback( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + ignore_unexpected: bool = False, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + udp_sock: Optional[Any] = None, + tcp_sock: Optional[Any] = None, +) -> Tuple[dns.message.Message, bool]: """Return the response to the query, trying UDP first and falling back to TCP if UDP results in a truncated response. @@ -635,26 +717,46 @@ def udp_with_fallback(q: dns.message.Message, where: str, timeout: Optional[floa if and only if TCP was used. """ try: - response = udp(q, where, timeout, port, source, source_port, - ignore_unexpected, one_rr_per_rrset, - ignore_trailing, True, udp_sock) + response = udp( + q, + where, + timeout, + port, + source, + source_port, + ignore_unexpected, + one_rr_per_rrset, + ignore_trailing, + True, + udp_sock, + ) return (response, False) except dns.message.Truncated: - response = tcp(q, where, timeout, port, source, source_port, - one_rr_per_rrset, ignore_trailing, tcp_sock) + response = tcp( + q, + where, + timeout, + port, + source, + source_port, + one_rr_per_rrset, + ignore_trailing, + tcp_sock, + ) return (response, True) + def _net_read(sock, count, expiration): """Read the specified number of bytes from sock. Keep trying until we either get the desired amount, or we hit EOF. A Timeout exception will be raised if the operation is not completed by the expiration time. """ - s = b'' + s = b"" while count > 0: try: n = sock.recv(count) - if n == b'': + if n == b"": raise EOFError count -= len(n) s += n @@ -681,8 +783,11 @@ def _net_write(sock, data, expiration): _wait_for_readable(sock, expiration) -def send_tcp(sock: Any, what: Union[dns.message.Message, bytes], - expiration: Optional[float]=None) -> Tuple[int, float]: +def send_tcp( + sock: Any, + what: Union[dns.message.Message, bytes], + expiration: Optional[float] = None, +) -> Tuple[int, float]: """Send a DNS message to the specified TCP socket. *sock*, a ``socket``. @@ -709,10 +814,15 @@ def send_tcp(sock: Any, what: Union[dns.message.Message, bytes], _net_write(sock, tcpmsg, expiration) return (len(tcpmsg), sent_time) -def receive_tcp(sock: Any, expiration: Optional[float]=None, one_rr_per_rrset: bool=False, - keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]]=None, - request_mac: Optional[bytes]=b'', - ignore_trailing: bool=False) -> Tuple[dns.message.Message, float]: + +def receive_tcp( + sock: Any, + expiration: Optional[float] = None, + one_rr_per_rrset: bool = False, + keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]] = None, + request_mac: Optional[bytes] = b"", + ignore_trailing: bool = False, +) -> Tuple[dns.message.Message, float]: """Read a DNS message from a TCP socket. *sock*, a ``socket``. @@ -742,11 +852,16 @@ def receive_tcp(sock: Any, expiration: Optional[float]=None, one_rr_per_rrset: b (l,) = struct.unpack("!H", ldata) wire = _net_read(sock, l, expiration) received_time = time.time() - r = dns.message.from_wire(wire, keyring=keyring, request_mac=request_mac, - one_rr_per_rrset=one_rr_per_rrset, - ignore_trailing=ignore_trailing) + r = dns.message.from_wire( + wire, + keyring=keyring, + request_mac=request_mac, + one_rr_per_rrset=one_rr_per_rrset, + ignore_trailing=ignore_trailing, + ) return (r, received_time) + def _connect(s, address, expiration): err = s.connect_ex(address) if err == 0: @@ -758,10 +873,17 @@ def _connect(s, address, expiration): raise OSError(err, os.strerror(err)) -def tcp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: int=53, - source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, - sock: Optional[Any]=None) -> dns.message.Message: +def tcp( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 53, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + sock: Optional[Any] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via TCP. *q*, a ``dns.message.Message``, the query to send @@ -800,20 +922,22 @@ def tcp(q: dns.message.Message, where: str, timeout: Optional[float]=None, port: if sock: s = sock else: - (af, destination, source) = _destination_and_source(where, port, - source, - source_port) - s = stack.enter_context(_make_socket(af, socket.SOCK_STREAM, - source)) + (af, destination, source) = _destination_and_source( + where, port, source, source_port + ) + s = stack.enter_context(_make_socket(af, socket.SOCK_STREAM, source)) _connect(s, destination, expiration) send_tcp(s, wire, expiration) - (r, received_time) = receive_tcp(s, expiration, one_rr_per_rrset, - q.keyring, q.mac, ignore_trailing) + (r, received_time) = receive_tcp( + s, expiration, one_rr_per_rrset, q.keyring, q.mac, ignore_trailing + ) r.time = received_time - begin_time if not q.is_response(r): raise BadResponse return r - assert False # help mypy figure out we can't get here lgtm[py/unreachable-statement] + assert ( + False # help mypy figure out we can't get here lgtm[py/unreachable-statement] + ) def _tls_handshake(s, expiration): @@ -827,11 +951,19 @@ def _tls_handshake(s, expiration): _wait_for_writable(s, expiration) -def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, - port: int=853, source: Optional[str]=None, source_port: int=0, - one_rr_per_rrset: bool=False, ignore_trailing: bool=False, sock: Optional[ssl.SSLSocket]=None, - ssl_context: Optional[ssl.SSLContext]=None, - server_hostname: Optional[str]=None) -> dns.message.Message: +def tls( + q: dns.message.Message, + where: str, + timeout: Optional[float] = None, + port: int = 853, + source: Optional[str] = None, + source_port: int = 0, + one_rr_per_rrset: bool = False, + ignore_trailing: bool = False, + sock: Optional[ssl.SSLSocket] = None, + ssl_context: Optional[ssl.SSLContext] = None, + server_hostname: Optional[str] = None, +) -> dns.message.Message: """Return the response obtained after sending a query via TLS. *q*, a ``dns.message.Message``, the query to send @@ -878,13 +1010,23 @@ def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, # # If a socket was provided, there's no special TLS handling needed. # - return tcp(q, where, timeout, port, source, source_port, - one_rr_per_rrset, ignore_trailing, sock) + return tcp( + q, + where, + timeout, + port, + source, + source_port, + one_rr_per_rrset, + ignore_trailing, + sock, + ) wire = q.to_wire() (begin_time, expiration) = _compute_times(timeout) - (af, destination, source) = _destination_and_source(where, port, - source, source_port) + (af, destination, source) = _destination_and_source( + where, port, source, source_port + ) if ssl_context is None and not sock: # LGTM complains about this because the default might permit TLS < 1.2 # for compatibility, but the python documentation says that explicit @@ -897,28 +1039,45 @@ def tls(q: dns.message.Message, where: str, timeout: Optional[float]=None, if server_hostname is None: ssl_context.check_hostname = False - with _make_socket(af, socket.SOCK_STREAM, source, ssl_context=ssl_context, - server_hostname=server_hostname) as s: + with _make_socket( + af, + socket.SOCK_STREAM, + source, + ssl_context=ssl_context, + server_hostname=server_hostname, + ) as s: _connect(s, destination, expiration) _tls_handshake(s, expiration) send_tcp(s, wire, expiration) - (r, received_time) = receive_tcp(s, expiration, one_rr_per_rrset, - q.keyring, q.mac, ignore_trailing) + (r, received_time) = receive_tcp( + s, expiration, one_rr_per_rrset, q.keyring, q.mac, ignore_trailing + ) r.time = received_time - begin_time if not q.is_response(r): raise BadResponse return r - assert False # help mypy figure out we can't get here lgtm[py/unreachable-statement] - -def xfr(where: str, zone: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.AXFR, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - timeout: Optional[float]=None, port: int=53, - keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]]=None, - keyname: Optional[Union[dns.name.Name, str]]=None, relativize: bool=True, - lifetime: Optional[float]=None, source: Optional[str]=None, source_port: int=0, - serial: int=0, use_udp: bool=False, - keyalgorithm: Union[dns.name.Name, str]=dns.tsig.default_algorithm) -> Any: + assert ( + False # help mypy figure out we can't get here lgtm[py/unreachable-statement] + ) + + +def xfr( + where: str, + zone: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.AXFR, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + timeout: Optional[float] = None, + port: int = 53, + keyring: Optional[Dict[dns.name.Name, dns.tsig.Key]] = None, + keyname: Optional[Union[dns.name.Name, str]] = None, + relativize: bool = True, + lifetime: Optional[float] = None, + source: Optional[str] = None, + source_port: int = 0, + serial: int = 0, + use_udp: bool = False, + keyalgorithm: Union[dns.name.Name, str] = dns.tsig.default_algorithm, +) -> Any: """Return a generator for the responses to a zone transfer. *where*, a ``str`` containing an IPv4 or IPv6 address, where @@ -976,16 +1135,16 @@ def xfr(where: str, zone: Union[dns.name.Name, str], rdtype = dns.rdatatype.RdataType.make(rdtype) q = dns.message.make_query(zone, rdtype, rdclass) if rdtype == dns.rdatatype.IXFR: - rrset = dns.rrset.from_text(zone, 0, 'IN', 'SOA', - '. . %u 0 0 0 0' % serial) + rrset = dns.rrset.from_text(zone, 0, "IN", "SOA", ". . %u 0 0 0 0" % serial) q.authority.append(rrset) if keyring is not None: q.use_tsig(keyring, keyname, algorithm=keyalgorithm) wire = q.to_wire() - (af, destination, source) = _destination_and_source(where, port, - source, source_port) + (af, destination, source) = _destination_and_source( + where, port, source, source_port + ) if use_udp and rdtype != dns.rdatatype.IXFR: - raise ValueError('cannot do a UDP AXFR') + raise ValueError("cannot do a UDP AXFR") sock_type = socket.SOCK_DGRAM if use_udp else socket.SOCK_STREAM with _make_socket(af, sock_type, source) as s: (_, expiration) = _compute_times(lifetime) @@ -1009,8 +1168,9 @@ def xfr(where: str, zone: Union[dns.name.Name, str], tsig_ctx = None while not done: (_, mexpiration) = _compute_times(timeout) - if mexpiration is None or \ - (expiration is not None and mexpiration > expiration): + if mexpiration is None or ( + expiration is not None and mexpiration > expiration + ): mexpiration = expiration if use_udp: (wire, _) = _udp_recv(s, 65535, mexpiration) @@ -1018,11 +1178,17 @@ def xfr(where: str, zone: Union[dns.name.Name, str], ldata = _net_read(s, 2, mexpiration) (l,) = struct.unpack("!H", ldata) wire = _net_read(s, l, mexpiration) - is_ixfr = (rdtype == dns.rdatatype.IXFR) - r = dns.message.from_wire(wire, keyring=q.keyring, - request_mac=q.mac, xfr=True, - origin=origin, tsig_ctx=tsig_ctx, - multi=True, one_rr_per_rrset=is_ixfr) + is_ixfr = rdtype == dns.rdatatype.IXFR + r = dns.message.from_wire( + wire, + keyring=q.keyring, + request_mac=q.mac, + xfr=True, + origin=origin, + tsig_ctx=tsig_ctx, + multi=True, + one_rr_per_rrset=is_ixfr, + ) rcode = r.rcode() if rcode != dns.rcode.NOERROR: raise TransferError(rcode) @@ -1030,8 +1196,7 @@ def xfr(where: str, zone: Union[dns.name.Name, str], answer_index = 0 if soa_rrset is None: if not r.answer or r.answer[0].name != oname: - raise dns.exception.FormError( - "No answer or RRset not for qname") + raise dns.exception.FormError("No answer or RRset not for qname") rrset = r.answer[0] if rrset.rdtype != dns.rdatatype.SOA: raise dns.exception.FormError("first RRset is not an SOA") @@ -1055,8 +1220,7 @@ def xfr(where: str, zone: Union[dns.name.Name, str], if rrset.rdtype == dns.rdatatype.SOA and rrset.name == oname: if expecting_SOA: if rrset[0].serial != serial: - raise dns.exception.FormError( - "IXFR base serial mismatch") + raise dns.exception.FormError("IXFR base serial mismatch") expecting_SOA = False elif rdtype == dns.rdatatype.IXFR: delete_mode = not delete_mode @@ -1065,9 +1229,10 @@ def xfr(where: str, zone: Union[dns.name.Name, str], # finished. If this is an IXFR we also check that we're # seeing the record in the expected part of the response. # - if rrset == soa_rrset and \ - (rdtype == dns.rdatatype.AXFR or - (rdtype == dns.rdatatype.IXFR and delete_mode)): + if rrset == soa_rrset and ( + rdtype == dns.rdatatype.AXFR + or (rdtype == dns.rdatatype.IXFR and delete_mode) + ): done = True elif expecting_SOA: # @@ -1089,15 +1254,23 @@ class UDPMode(enum.IntEnum): TRY_FIRST means "try to use UDP but fall back to TCP if needed" ONLY means "raise ``dns.xfr.UseTCP`` if trying UDP does not succeed" """ + NEVER = 0 TRY_FIRST = 1 ONLY = 2 -def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManager, - query: Optional[dns.message.Message]=None, - port: int=53, timeout: Optional[float]=None, lifetime: Optional[float]=None, - source: Optional[str]=None, source_port: int=0, udp_mode: UDPMode=UDPMode.NEVER) -> None: +def inbound_xfr( + where: str, + txn_manager: dns.transaction.TransactionManager, + query: Optional[dns.message.Message] = None, + port: int = 53, + timeout: Optional[float] = None, + lifetime: Optional[float] = None, + source: Optional[str] = None, + source_port: int = 0, + udp_mode: UDPMode = UDPMode.NEVER, +) -> None: """Conduct an inbound transfer and apply it via a transaction from the txn_manager. @@ -1142,8 +1315,9 @@ def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManager, is_ixfr = rdtype == dns.rdatatype.IXFR origin = txn_manager.from_wire_origin() wire = query.to_wire() - (af, destination, source) = _destination_and_source(where, port, - source, source_port) + (af, destination, source) = _destination_and_source( + where, port, source, source_port + ) (_, expiration) = _compute_times(lifetime) retry = True while retry: @@ -1161,14 +1335,14 @@ def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManager, else: tcpmsg = struct.pack("!H", len(wire)) + wire _net_write(s, tcpmsg, expiration) - with dns.xfr.Inbound(txn_manager, rdtype, serial, - is_udp) as inbound: + with dns.xfr.Inbound(txn_manager, rdtype, serial, is_udp) as inbound: done = False tsig_ctx = None while not done: (_, mexpiration) = _compute_times(timeout) - if mexpiration is None or \ - (expiration is not None and mexpiration > expiration): + if mexpiration is None or ( + expiration is not None and mexpiration > expiration + ): mexpiration = expiration if is_udp: (rwire, _) = _udp_recv(s, 65535, mexpiration) @@ -1176,11 +1350,16 @@ def inbound_xfr(where: str, txn_manager: dns.transaction.TransactionManager, ldata = _net_read(s, 2, mexpiration) (l,) = struct.unpack("!H", ldata) rwire = _net_read(s, l, mexpiration) - r = dns.message.from_wire(rwire, keyring=query.keyring, - request_mac=query.mac, xfr=True, - origin=origin, tsig_ctx=tsig_ctx, - multi=(not is_udp), - one_rr_per_rrset=is_ixfr) + r = dns.message.from_wire( + rwire, + keyring=query.keyring, + request_mac=query.mac, + xfr=True, + origin=origin, + tsig_ctx=tsig_ctx, + multi=(not is_udp), + one_rr_per_rrset=is_ixfr, + ) try: done = inbound.process_message(r) except dns.xfr.UseTCP: diff --git a/dns/rcode.py b/dns/rcode.py index 16e1ed4..8e6386f 100644 --- a/dns/rcode.py +++ b/dns/rcode.py @@ -22,6 +22,7 @@ from typing import Tuple import dns.enum import dns.exception + class Rcode(dns.enum.IntEnum): #: No error NOERROR = 0 @@ -104,7 +105,7 @@ def from_flags(flags: int, ednsflags: int) -> Rcode: Returns a ``dns.rcode.Rcode``. """ - value = (flags & 0x000f) | ((ednsflags >> 20) & 0xff0) + value = (flags & 0x000F) | ((ednsflags >> 20) & 0xFF0) return Rcode.make(value) @@ -119,13 +120,13 @@ def to_flags(value: Rcode) -> Tuple[int, int]: """ if value < 0 or value > 4095: - raise ValueError('rcode must be >= 0 and <= 4095') - v = value & 0xf - ev = (value & 0xff0) << 20 + raise ValueError("rcode must be >= 0 and <= 4095") + v = value & 0xF + ev = (value & 0xFF0) << 20 return (v, ev) -def to_text(value: Rcode, tsig: bool=False) -> str: +def to_text(value: Rcode, tsig: bool = False) -> str: """Convert rcode into text. *value*, a ``dns.rcode.Rcode``, the rcode. @@ -136,9 +137,10 @@ def to_text(value: Rcode, tsig: bool=False) -> str: """ if tsig and value == Rcode.BADVERS: - return 'BADSIG' + return "BADSIG" return Rcode.to_text(value) + ### BEGIN generated Rcode constants NOERROR = Rcode.NOERROR diff --git a/dns/rdata.py b/dns/rdata.py index 155e124..dc2ad97 100644 --- a/dns/rdata.py +++ b/dns/rdata.py @@ -57,21 +57,22 @@ class NoRelativeRdataOrdering(dns.exception.DNSException): """ -def _wordbreak(data, chunksize=_chunksize, separator=b' '): +def _wordbreak(data, chunksize=_chunksize, separator=b" "): """Break a binary string into chunks of chunksize characters separated by a space. """ if not chunksize: return data.decode() - return separator.join([data[i:i + chunksize] - for i - in range(0, len(data), chunksize)]).decode() + return separator.join( + [data[i : i + chunksize] for i in range(0, len(data), chunksize)] + ).decode() # pylint: disable=unused-argument -def _hexify(data, chunksize=_chunksize, separator=b' ', **kw): + +def _hexify(data, chunksize=_chunksize, separator=b" ", **kw): """Convert a binary string into its hex encoding, broken up into chunks of chunksize characters separated by a separator. """ @@ -79,17 +80,19 @@ def _hexify(data, chunksize=_chunksize, separator=b' ', **kw): return _wordbreak(binascii.hexlify(data), chunksize, separator) -def _base64ify(data, chunksize=_chunksize, separator=b' ', **kw): +def _base64ify(data, chunksize=_chunksize, separator=b" ", **kw): """Convert a binary string into its base64 encoding, broken up into chunks of chunksize characters separated by a separator. """ return _wordbreak(base64.b64encode(data), chunksize, separator) + # pylint: enable=unused-argument __escaped = b'"\\' + def _escapify(qstring): """Escape the characters in a quoted string which need it.""" @@ -98,14 +101,14 @@ def _escapify(qstring): if not isinstance(qstring, bytearray): qstring = bytearray(qstring) - text = '' + text = "" for c in qstring: if c in __escaped: - text += '\\' + chr(c) + text += "\\" + chr(c) elif c >= 0x20 and c < 0x7F: text += chr(c) else: - text += '\\%03d' % c + text += "\\%03d" % c return text @@ -116,9 +119,10 @@ def _truncate_bitmap(what): for i in range(len(what) - 1, -1, -1): if what[i] != 0: - return what[0: i + 1] + return what[0 : i + 1] return what[0:1] + # So we don't have to edit all the rdata classes... _constify = dns.immutable.constify @@ -127,7 +131,7 @@ _constify = dns.immutable.constify class Rdata: """Base class for all DNS rdata types.""" - __slots__ = ['rdclass', 'rdtype', 'rdcomment'] + __slots__ = ["rdclass", "rdtype", "rdcomment"] def __init__(self, rdclass, rdtype): """Initialize an rdata. @@ -142,8 +146,9 @@ class Rdata: self.rdcomment: Optional[str] = None def _get_all_slots(self): - return itertools.chain.from_iterable(getattr(cls, '__slots__', []) - for cls in self.__class__.__mro__) + return itertools.chain.from_iterable( + getattr(cls, "__slots__", []) for cls in self.__class__.__mro__ + ) def __getstate__(self): # We used to try to do a tuple of all slots here, but it @@ -162,10 +167,10 @@ class Rdata: def __setstate__(self, state): for slot, val in state.items(): object.__setattr__(self, slot, val) - if not hasattr(self, 'rdcomment'): + if not hasattr(self, "rdcomment"): # Pickled rdata from 2.0.x might not have a rdcomment, so add # it if needed. - object.__setattr__(self, 'rdcomment', None) + object.__setattr__(self, "rdcomment", None) def covers(self) -> dns.rdatatype.RdataType: """Return the type a Rdata covers. @@ -191,7 +196,12 @@ class Rdata: return self.covers() << 16 | self.rdtype - def to_text(self, origin: Optional[dns.name.Name]=None, relativize: bool=True, **kw: Dict[str, Any]) -> str: + def to_text( + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + **kw: Dict[str, Any] + ) -> str: """Convert an rdata to text format. Returns a ``str``. @@ -199,12 +209,22 @@ class Rdata: raise NotImplementedError # pragma: no cover - def _to_wire(self, file: Optional[Any], compress: Optional[dns.name.CompressType]=None, - origin: Optional[dns.name.Name]=None, canonicalize: bool=False) -> bytes: + def _to_wire( + self, + file: Optional[Any], + compress: Optional[dns.name.CompressType] = None, + origin: Optional[dns.name.Name] = None, + canonicalize: bool = False, + ) -> bytes: raise NotImplementedError # pragma: no cover - def to_wire(self, file: Optional[Any]=None, compress: Optional[dns.name.CompressType]=None, - origin: Optional[dns.name.Name]=None, canonicalize: bool=False) -> bytes: + def to_wire( + self, + file: Optional[Any] = None, + compress: Optional[dns.name.CompressType] = None, + origin: Optional[dns.name.Name] = None, + canonicalize: bool = False, + ) -> bytes: """Convert an rdata to wire format. Returns a ``bytes`` or ``None``. @@ -217,15 +237,18 @@ class Rdata: self._to_wire(f, compress, origin, canonicalize) return f.getvalue() - def to_generic(self, origin: Optional[dns.name.Name]=None) -> 'dns.rdata.GenericRdata': + def to_generic( + self, origin: Optional[dns.name.Name] = None + ) -> "dns.rdata.GenericRdata": """Creates a dns.rdata.GenericRdata equivalent of this rdata. Returns a ``dns.rdata.GenericRdata``. """ - return dns.rdata.GenericRdata(self.rdclass, self.rdtype, - self.to_wire(origin=origin)) + return dns.rdata.GenericRdata( + self.rdclass, self.rdtype, self.to_wire(origin=origin) + ) - def to_digestable(self, origin: Optional[dns.name.Name]=None) -> bytes: + def to_digestable(self, origin: Optional[dns.name.Name] = None) -> bytes: """Convert rdata to a format suitable for digesting in hashes. This is also the DNSSEC canonical form. @@ -237,12 +260,19 @@ class Rdata: def __repr__(self): covers = self.covers() if covers == dns.rdatatype.NONE: - ctext = '' + ctext = "" else: - ctext = '(' + dns.rdatatype.to_text(covers) + ')' - return '<DNS ' + dns.rdataclass.to_text(self.rdclass) + ' ' + \ - dns.rdatatype.to_text(self.rdtype) + ctext + ' rdata: ' + \ - str(self) + '>' + ctext = "(" + dns.rdatatype.to_text(covers) + ")" + return ( + "<DNS " + + dns.rdataclass.to_text(self.rdclass) + + " " + + dns.rdatatype.to_text(self.rdtype) + + ctext + + " rdata: " + + str(self) + + ">" + ) def __str__(self): return self.to_text() @@ -323,27 +353,39 @@ class Rdata: return not self.__eq__(other) def __lt__(self, other): - if not isinstance(other, Rdata) or \ - self.rdclass != other.rdclass or self.rdtype != other.rdtype: + if ( + not isinstance(other, Rdata) + or self.rdclass != other.rdclass + or self.rdtype != other.rdtype + ): return NotImplemented return self._cmp(other) < 0 def __le__(self, other): - if not isinstance(other, Rdata) or \ - self.rdclass != other.rdclass or self.rdtype != other.rdtype: + if ( + not isinstance(other, Rdata) + or self.rdclass != other.rdclass + or self.rdtype != other.rdtype + ): return NotImplemented return self._cmp(other) <= 0 def __ge__(self, other): - if not isinstance(other, Rdata) or \ - self.rdclass != other.rdclass or self.rdtype != other.rdtype: + if ( + not isinstance(other, Rdata) + or self.rdclass != other.rdclass + or self.rdtype != other.rdtype + ): return NotImplemented return self._cmp(other) >= 0 def __gt__(self, other): - if not isinstance(other, Rdata) or \ - self.rdclass != other.rdclass or self.rdtype != other.rdtype: + if ( + not isinstance(other, Rdata) + or self.rdclass != other.rdclass + or self.rdtype != other.rdtype + ): return NotImplemented return self._cmp(other) > 0 @@ -351,19 +393,28 @@ class Rdata: return hash(self.to_digestable(dns.name.root)) @classmethod - def from_text(cls, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - tok: dns.tokenizer.Tokenizer, origin: Optional[dns.name.Name]=None, relativize: bool=True, - relativize_to: Optional[dns.name.Name]=None) -> 'Rdata': + def from_text( + cls, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + tok: dns.tokenizer.Tokenizer, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, + ) -> "Rdata": raise NotImplementedError # pragma: no cover @classmethod - def from_wire_parser(cls, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - parser: dns.wire.Parser, origin: Optional[dns.name.Name]=None) -> 'Rdata': + def from_wire_parser( + cls, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + parser: dns.wire.Parser, + origin: Optional[dns.name.Name] = None, + ) -> "Rdata": raise NotImplementedError # pragma: no cover - def replace(self, **kwargs: Dict[str, Any]) -> 'Rdata': + def replace(self, **kwargs: Dict[str, Any]) -> "Rdata": """ Create a new Rdata instance based on the instance replace was invoked on. It is possible to pass different parameters to @@ -381,14 +432,20 @@ class Rdata: # Ensure that all of the arguments correspond to valid fields. # Don't allow rdclass or rdtype to be changed, though. for key in kwargs: - if key == 'rdcomment': + if key == "rdcomment": continue if key not in parameters: - raise AttributeError("'{}' object has no attribute '{}'" - .format(self.__class__.__name__, key)) - if key in ('rdclass', 'rdtype'): - raise AttributeError("Cannot overwrite '{}' attribute '{}'" - .format(self.__class__.__name__, key)) + raise AttributeError( + "'{}' object has no attribute '{}'".format( + self.__class__.__name__, key + ) + ) + if key in ("rdclass", "rdtype"): + raise AttributeError( + "Cannot overwrite '{}' attribute '{}'".format( + self.__class__.__name__, key + ) + ) # Construct the parameter list. For each field, use the value in # kwargs if present, and the current value otherwise. @@ -398,9 +455,9 @@ class Rdata: rd = self.__class__(*args) # The comment is not set in the constructor, so give it special # handling. - rdcomment = kwargs.get('rdcomment', self.rdcomment) + rdcomment = kwargs.get("rdcomment", self.rdcomment) if rdcomment is not None: - object.__setattr__(rd, 'rdcomment', rdcomment) + object.__setattr__(rd, "rdcomment", rdcomment) return rd # Type checking and conversion helpers. These are class methods as @@ -415,8 +472,13 @@ class Rdata: return dns.rdatatype.RdataType.make(value) @classmethod - def _as_bytes(cls, value: Any, encode: bool=False, max_length: Optional[int]=None, - empty_ok: bool=True) -> bytes: + def _as_bytes( + cls, + value: Any, + encode: bool = False, + max_length: Optional[int] = None, + empty_ok: bool = True, + ) -> bytes: if encode and isinstance(value, str): bvalue = value.encode() elif isinstance(value, bytearray): @@ -424,11 +486,11 @@ class Rdata: elif isinstance(value, bytes): bvalue = value else: - raise ValueError('not bytes') + raise ValueError("not bytes") if max_length is not None and len(bvalue) > max_length: - raise ValueError('too long') + raise ValueError("too long") if not empty_ok and len(bvalue) == 0: - raise ValueError('empty bytes not allowed') + raise ValueError("empty bytes not allowed") return bvalue @classmethod @@ -439,49 +501,49 @@ class Rdata: if isinstance(value, str): return dns.name.from_text(value) elif not isinstance(value, dns.name.Name): - raise ValueError('not a name') + raise ValueError("not a name") return value @classmethod def _as_uint8(cls, value): if not isinstance(value, int): - raise ValueError('not an integer') + raise ValueError("not an integer") if value < 0 or value > 255: - raise ValueError('not a uint8') + raise ValueError("not a uint8") return value @classmethod def _as_uint16(cls, value): if not isinstance(value, int): - raise ValueError('not an integer') + raise ValueError("not an integer") if value < 0 or value > 65535: - raise ValueError('not a uint16') + raise ValueError("not a uint16") return value @classmethod def _as_uint32(cls, value): if not isinstance(value, int): - raise ValueError('not an integer') + raise ValueError("not an integer") if value < 0 or value > 4294967295: - raise ValueError('not a uint32') + raise ValueError("not a uint32") return value @classmethod def _as_uint48(cls, value): if not isinstance(value, int): - raise ValueError('not an integer') + raise ValueError("not an integer") if value < 0 or value > 281474976710655: - raise ValueError('not a uint48') + raise ValueError("not a uint48") return value @classmethod def _as_int(cls, value, low=None, high=None): if not isinstance(value, int): - raise ValueError('not an integer') + raise ValueError("not an integer") if low is not None and value < low: - raise ValueError('value too small') + raise ValueError("value too small") if high is not None and value > high: - raise ValueError('value too large') + raise ValueError("value too large") return value @classmethod @@ -493,7 +555,7 @@ class Rdata: elif isinstance(value, bytes): return dns.ipv4.inet_ntoa(value) else: - raise ValueError('not an IPv4 address') + raise ValueError("not an IPv4 address") @classmethod def _as_ipv6_address(cls, value): @@ -504,14 +566,14 @@ class Rdata: elif isinstance(value, bytes): return dns.ipv6.inet_ntoa(value) else: - raise ValueError('not an IPv6 address') + raise ValueError("not an IPv6 address") @classmethod def _as_bool(cls, value): if isinstance(value, bool): return value else: - raise ValueError('not a boolean') + raise ValueError("not a boolean") @classmethod def _as_ttl(cls, value): @@ -520,7 +582,7 @@ class Rdata: elif isinstance(value, str): return dns.ttl.from_text(value) else: - raise ValueError('not a TTL') + raise ValueError("not a TTL") @classmethod def _as_tuple(cls, value, as_value): @@ -541,6 +603,7 @@ class Rdata: random.shuffle(items) return items + @dns.immutable.immutable class GenericRdata(Rdata): @@ -550,28 +613,32 @@ class GenericRdata(Rdata): implementation. It implements the DNS "unknown RRs" scheme. """ - __slots__ = ['data'] + __slots__ = ["data"] def __init__(self, rdclass, rdtype, data): super().__init__(rdclass, rdtype) self.data = data - def to_text(self, origin: Optional[dns.name.Name]=None, relativize: bool=True, **kw: Dict[str, Any]) -> str: - return r'\# %d ' % len(self.data) + _hexify(self.data, **kw) + def to_text( + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + **kw: Dict[str, Any] + ) -> str: + return r"\# %d " % len(self.data) + _hexify(self.data, **kw) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): token = tok.get() - if not token.is_identifier() or token.value != r'\#': - raise dns.exception.SyntaxError( - r'generic rdata does not start with \#') + if not token.is_identifier() or token.value != r"\#": + raise dns.exception.SyntaxError(r"generic rdata does not start with \#") length = tok.get_int() hex = tok.concatenate_remaining_identifiers(True).encode() data = binascii.unhexlify(hex) if len(data) != length: - raise dns.exception.SyntaxError( - 'generic rdata hex data has wrong length') + raise dns.exception.SyntaxError("generic rdata hex data has wrong length") return cls(rdclass, rdtype, data) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): @@ -581,8 +648,12 @@ class GenericRdata(Rdata): def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): return cls(rdclass, rdtype, parser.get_remaining()) -_rdata_classes: Dict[Tuple[dns.rdataclass.RdataClass, dns.rdatatype.RdataType], Any] = {} -_module_prefix = 'dns.rdtypes' + +_rdata_classes: Dict[ + Tuple[dns.rdataclass.RdataClass, dns.rdatatype.RdataType], Any +] = {} +_module_prefix = "dns.rdtypes" + def get_rdata_class(rdclass, rdtype): cls = _rdata_classes.get((rdclass, rdtype)) @@ -591,16 +662,16 @@ def get_rdata_class(rdclass, rdtype): if not cls: rdclass_text = dns.rdataclass.to_text(rdclass) rdtype_text = dns.rdatatype.to_text(rdtype) - rdtype_text = rdtype_text.replace('-', '_') + rdtype_text = rdtype_text.replace("-", "_") try: - mod = import_module('.'.join([_module_prefix, - rdclass_text, rdtype_text])) + mod = import_module( + ".".join([_module_prefix, rdclass_text, rdtype_text]) + ) cls = getattr(mod, rdtype_text) _rdata_classes[(rdclass, rdtype)] = cls except ImportError: try: - mod = import_module('.'.join([_module_prefix, - 'ANY', rdtype_text])) + mod = import_module(".".join([_module_prefix, "ANY", rdtype_text])) cls = getattr(mod, rdtype_text) _rdata_classes[(dns.rdataclass.ANY, rdtype)] = cls _rdata_classes[(rdclass, rdtype)] = cls @@ -612,12 +683,15 @@ def get_rdata_class(rdclass, rdtype): return cls -def from_text(rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - tok: Union[dns.tokenizer.Tokenizer, str], - origin: Optional[dns.name.Name]=None, - relativize: bool=True, relativize_to: Optional[dns.name.Name]=None, - idna_codec: Optional[dns.name.IDNACodec]=None) -> Rdata: +def from_text( + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + tok: Union[dns.tokenizer.Tokenizer, str], + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, + idna_codec: Optional[dns.name.IDNACodec] = None, +) -> Rdata: """Build an rdata object from text format. This function attempts to dynamically load a class which @@ -665,17 +739,18 @@ def from_text(rdclass: Union[dns.rdataclass.RdataClass, str], # peek at first token token = tok.get() tok.unget(token) - if token.is_identifier() and \ - token.value == r'\#': + if token.is_identifier() and token.value == r"\#": # # Known type using the generic syntax. Extract the # wire form from the generic syntax, and then run # from_wire on it. # - grdata = GenericRdata.from_text(rdclass, rdtype, tok, origin, - relativize, relativize_to) - rdata = from_wire(rdclass, rdtype, grdata.data, 0, - len(grdata.data), origin) + grdata = GenericRdata.from_text( + rdclass, rdtype, tok, origin, relativize, relativize_to + ) + rdata = from_wire( + rdclass, rdtype, grdata.data, 0, len(grdata.data), origin + ) # # If this comparison isn't equal, then there must have been # compressed names in the wire format, which is an error, @@ -683,21 +758,27 @@ def from_text(rdclass: Union[dns.rdataclass.RdataClass, str], # rwire = rdata.to_wire() if rwire != grdata.data: - raise dns.exception.SyntaxError('compressed data in ' - 'generic syntax form ' - 'of known rdatatype') + raise dns.exception.SyntaxError( + "compressed data in " + "generic syntax form " + "of known rdatatype" + ) if rdata is None: - rdata = cls.from_text(rdclass, rdtype, tok, origin, relativize, - relativize_to) + rdata = cls.from_text( + rdclass, rdtype, tok, origin, relativize, relativize_to + ) token = tok.get_eol_as_token() if token.comment is not None: - object.__setattr__(rdata, 'rdcomment', token.comment) + object.__setattr__(rdata, "rdcomment", token.comment) return rdata -def from_wire_parser(rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - parser: dns.wire.Parser, origin: Optional[dns.name.Name]=None) -> Rdata: +def from_wire_parser( + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + parser: dns.wire.Parser, + origin: Optional[dns.name.Name] = None, +) -> Rdata: """Build an rdata object from wire format This function attempts to dynamically load a class which @@ -728,10 +809,14 @@ def from_wire_parser(rdclass: Union[dns.rdataclass.RdataClass, str], return cls.from_wire_parser(rdclass, rdtype, parser, origin) -def from_wire(rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - wire: bytes, current: int, rdlen: int, - origin: Optional[dns.name.Name]=None) -> Rdata: +def from_wire( + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + wire: bytes, + current: int, + rdlen: int, + origin: Optional[dns.name.Name] = None, +) -> Rdata: """Build an rdata object from wire format This function attempts to dynamically load a class which @@ -765,13 +850,21 @@ def from_wire(rdclass: Union[dns.rdataclass.RdataClass, str], class RdatatypeExists(dns.exception.DNSException): """DNS rdatatype already exists.""" - supp_kwargs = {'rdclass', 'rdtype'} - fmt = "The rdata type with class {rdclass:d} and rdtype {rdtype:d} " + \ - "already exists." + + supp_kwargs = {"rdclass", "rdtype"} + fmt = ( + "The rdata type with class {rdclass:d} and rdtype {rdtype:d} " + + "already exists." + ) -def register_type(implementation: Any, rdtype: int, rdtype_text: str, is_singleton: bool=False, - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN) -> None: +def register_type( + implementation: Any, + rdtype: int, + rdtype_text: str, + is_singleton: bool = False, + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, +) -> None: """Dynamically register a module to handle an rdatatype. *implementation*, a module implementing the type in the usual dnspython @@ -797,6 +890,7 @@ def register_type(implementation: Any, rdtype: int, rdtype_text: str, is_singlet raise RdatatypeExists(rdclass=rdclass, rdtype=the_rdtype) except ValueError: pass - _rdata_classes[(rdclass, the_rdtype)] = getattr(implementation, - rdtype_text.replace('-', '_')) + _rdata_classes[(rdclass, the_rdtype)] = getattr( + implementation, rdtype_text.replace("-", "_") + ) dns.rdatatype.register_type(the_rdtype, rdtype_text, is_singleton) diff --git a/dns/rdataclass.py b/dns/rdataclass.py index 2867054..89b85a7 100644 --- a/dns/rdataclass.py +++ b/dns/rdataclass.py @@ -20,8 +20,10 @@ import dns.enum import dns.exception + class RdataClass(dns.enum.IntEnum): """DNS Rdata Class""" + RESERVED0 = 0 IN = 1 INTERNET = IN @@ -100,6 +102,7 @@ def is_metaclass(rdclass: RdataClass) -> bool: return True return False + ### BEGIN generated RdataClass constants RESERVED0 = RdataClass.RESERVED0 diff --git a/dns/rdataset.py b/dns/rdataset.py index c4b8644..072e7f7 100644 --- a/dns/rdataset.py +++ b/dns/rdataset.py @@ -17,7 +17,7 @@ """DNS rdatasets (an rdataset is a set of rdatas of a given type and class)""" -from typing import Any, cast, Collection, Dict, Iterable, List, Optional, Union +from typing import Any, cast, Collection, Dict, List, Optional, Union import io import random @@ -49,11 +49,15 @@ class Rdataset(dns.set.Set): """A DNS rdataset.""" - __slots__ = ['rdclass', 'rdtype', 'covers', 'ttl'] + __slots__ = ["rdclass", "rdtype", "covers", "ttl"] - def __init__(self, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, ttl: int=0): + def __init__( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + ttl: int = 0, + ): """Create a new rdataset of the specified class and type. *rdclass*, a ``dns.rdataclass.RdataClass``, the rdataclass. @@ -94,7 +98,9 @@ class Rdataset(dns.set.Set): elif ttl < self.ttl: self.ttl = ttl - def add(self, rd: dns.rdata.Rdata, ttl: Optional[int]=None) -> None: # pylint: disable=arguments-differ + def add( + self, rd: dns.rdata.Rdata, ttl: Optional[int] = None + ) -> None: # pylint: disable=arguments-differ """Add the specified rdata to the rdataset. If the optional *ttl* parameter is supplied, then @@ -121,8 +127,7 @@ class Rdataset(dns.set.Set): raise IncompatibleTypes if ttl is not None: self.update_ttl(ttl) - if self.rdtype == dns.rdatatype.RRSIG or \ - self.rdtype == dns.rdatatype.SIG: + if self.rdtype == dns.rdatatype.RRSIG or self.rdtype == dns.rdatatype.SIG: covers = rd.covers() if len(self) == 0 and self.covers == dns.rdatatype.NONE: self.covers = covers @@ -153,19 +158,26 @@ class Rdataset(dns.set.Set): def _rdata_repr(self): def maybe_truncate(s): if len(s) > 100: - return s[:100] + '...' + return s[:100] + "..." return s - return '[%s]' % ', '.join('<%s>' % maybe_truncate(str(rr)) - for rr in self) + + return "[%s]" % ", ".join("<%s>" % maybe_truncate(str(rr)) for rr in self) def __repr__(self): if self.covers == 0: - ctext = '' + ctext = "" else: - ctext = '(' + dns.rdatatype.to_text(self.covers) + ')' - return '<DNS ' + dns.rdataclass.to_text(self.rdclass) + ' ' + \ - dns.rdatatype.to_text(self.rdtype) + ctext + \ - ' rdataset: ' + self._rdata_repr() + '>' + ctext = "(" + dns.rdatatype.to_text(self.covers) + ")" + return ( + "<DNS " + + dns.rdataclass.to_text(self.rdclass) + + " " + + dns.rdatatype.to_text(self.rdtype) + + ctext + + " rdataset: " + + self._rdata_repr() + + ">" + ) def __str__(self): return self.to_text() @@ -173,20 +185,26 @@ class Rdataset(dns.set.Set): def __eq__(self, other): if not isinstance(other, Rdataset): return False - if self.rdclass != other.rdclass or \ - self.rdtype != other.rdtype or \ - self.covers != other.covers: + if ( + self.rdclass != other.rdclass + or self.rdtype != other.rdtype + or self.covers != other.covers + ): return False return super().__eq__(other) def __ne__(self, other): return not self.__eq__(other) - def to_text(self, name: Optional[dns.name.Name]=None, - origin: Optional[dns.name.Name]=None, - relativize: bool=True, - override_rdclass: Optional[dns.rdataclass.RdataClass]=None, - want_comments: bool=False, **kw: Dict[str, Any]) -> str: + def to_text( + self, + name: Optional[dns.name.Name] = None, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + override_rdclass: Optional[dns.rdataclass.RdataClass] = None, + want_comments: bool = False, + **kw: Dict[str, Any], + ) -> str: """Convert the rdataset into DNS zone file format. See ``dns.name.Name.choose_relativity`` for more information @@ -215,10 +233,10 @@ class Rdataset(dns.set.Set): if name is not None: name = name.choose_relativity(origin, relativize) ntext = str(name) - pad = ' ' + pad = " " else: - ntext = '' - pad = '' + ntext = "" + pad = "" s = io.StringIO() if override_rdclass is not None: rdclass = override_rdclass @@ -230,31 +248,46 @@ class Rdataset(dns.set.Set): # some dynamic updates, so we don't need to print out the TTL # (which is meaningless anyway). # - s.write('{}{}{} {}\n'.format(ntext, pad, - dns.rdataclass.to_text(rdclass), - dns.rdatatype.to_text(self.rdtype))) + s.write( + "{}{}{} {}\n".format( + ntext, + pad, + dns.rdataclass.to_text(rdclass), + dns.rdatatype.to_text(self.rdtype), + ) + ) else: for rd in self: - extra = '' + extra = "" if want_comments: if rd.rdcomment: - extra = f' ;{rd.rdcomment}' - s.write('%s%s%d %s %s %s%s\n' % - (ntext, pad, self.ttl, dns.rdataclass.to_text(rdclass), - dns.rdatatype.to_text(self.rdtype), - rd.to_text(origin=origin, relativize=relativize, - **kw), - extra)) + extra = f" ;{rd.rdcomment}" + s.write( + "%s%s%d %s %s %s%s\n" + % ( + ntext, + pad, + self.ttl, + dns.rdataclass.to_text(rdclass), + dns.rdatatype.to_text(self.rdtype), + rd.to_text(origin=origin, relativize=relativize, **kw), + extra, + ) + ) # # We strip off the final \n for the caller's convenience in printing # return s.getvalue()[:-1] - def to_wire(self, name: dns.name.Name, file: Any, - compress: Optional[dns.name.CompressType]=None, - origin: Optional[dns.name.Name]=None, - override_rdclass: Optional[dns.rdataclass.RdataClass]=None, - want_shuffle: bool=True) -> int: + def to_wire( + self, + name: dns.name.Name, + file: Any, + compress: Optional[dns.name.CompressType] = None, + origin: Optional[dns.name.Name] = None, + override_rdclass: Optional[dns.rdataclass.RdataClass] = None, + want_shuffle: bool = True, + ) -> int: """Convert the rdataset to wire format. *name*, a ``dns.name.Name`` is the owner name to use. @@ -299,8 +332,7 @@ class Rdataset(dns.set.Set): l = self for rd in l: name.to_wire(file, compress, origin) - stuff = struct.pack("!HHIH", self.rdtype, rdclass, - self.ttl, 0) + stuff = struct.pack("!HHIH", self.rdtype, rdclass, self.ttl, 0) file.write(stuff) start = file.tell() rd.to_wire(file, compress, origin) @@ -312,15 +344,16 @@ class Rdataset(dns.set.Set): file.seek(0, io.SEEK_END) return len(self) - def match(self, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType) -> bool: + def match( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType, + ) -> bool: """Returns ``True`` if this rdataset matches the specified class, type, and covers. """ - if self.rdclass == rdclass and \ - self.rdtype == rdtype and \ - self.covers == covers: + if self.rdclass == rdclass and self.rdtype == rdtype and self.covers == covers: return True return False @@ -349,46 +382,47 @@ class ImmutableRdataset(Rdataset): # lgtm[py/missing-equals] def __init__(self, rdataset: Rdataset): """Create an immutable rdataset from the specified rdataset.""" - super().__init__(rdataset.rdclass, rdataset.rdtype, rdataset.covers, - rdataset.ttl) + super().__init__( + rdataset.rdclass, rdataset.rdtype, rdataset.covers, rdataset.ttl + ) self.items = dns.immutable.Dict(rdataset.items) def update_ttl(self, ttl): - raise TypeError('immutable') + raise TypeError("immutable") def add(self, rd, ttl=None): - raise TypeError('immutable') + raise TypeError("immutable") def union_update(self, other): - raise TypeError('immutable') + raise TypeError("immutable") def intersection_update(self, other): - raise TypeError('immutable') + raise TypeError("immutable") def update(self, other): - raise TypeError('immutable') + raise TypeError("immutable") def __delitem__(self, i): - raise TypeError('immutable') + raise TypeError("immutable") # lgtm complains about these not raising ArithmeticError, but there is # precedent for overrides of these methods in other classes to raise # TypeError, and it seems like the better exception. def __ior__(self, other): # lgtm[py/unexpected-raise-in-special-method] - raise TypeError('immutable') + raise TypeError("immutable") def __iand__(self, other): # lgtm[py/unexpected-raise-in-special-method] - raise TypeError('immutable') + raise TypeError("immutable") def __iadd__(self, other): # lgtm[py/unexpected-raise-in-special-method] - raise TypeError('immutable') + raise TypeError("immutable") def __isub__(self, other): # lgtm[py/unexpected-raise-in-special-method] - raise TypeError('immutable') + raise TypeError("immutable") def clear(self): - raise TypeError('immutable') + raise TypeError("immutable") def __copy__(self): return ImmutableRdataset(super().copy()) @@ -409,12 +443,16 @@ class ImmutableRdataset(Rdataset): # lgtm[py/missing-equals] return ImmutableRdataset(super().symmetric_difference(other)) -def from_text_list(rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - ttl: int, text_rdatas: Collection[str], - idna_codec: Optional[dns.name.IDNACodec]=None, - origin: Optional[dns.name.Name]=None, - relativize: bool=True, relativize_to: Optional[dns.name.Name]=None) -> Rdataset: +def from_text_list( + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + ttl: int, + text_rdatas: Collection[str], + idna_codec: Optional[dns.name.IDNACodec] = None, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, +) -> Rdataset: """Create an rdataset with the specified class, type, and TTL, and with the specified list of rdatas in text format. @@ -438,15 +476,19 @@ def from_text_list(rdclass: Union[dns.rdataclass.RdataClass, str], r = Rdataset(the_rdclass, the_rdtype) r.update_ttl(ttl) for t in text_rdatas: - rd = dns.rdata.from_text(r.rdclass, r.rdtype, t, origin, relativize, - relativize_to, idna_codec) + rd = dns.rdata.from_text( + r.rdclass, r.rdtype, t, origin, relativize, relativize_to, idna_codec + ) r.add(rd) return r -def from_text(rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - ttl: int, *text_rdatas: Any) -> Rdataset: +def from_text( + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + ttl: int, + *text_rdatas: Any, +) -> Rdataset: """Create an rdataset with the specified class, type, and TTL, and with the specified rdatas in text format. diff --git a/dns/rdatatype.py b/dns/rdatatype.py index aded5bd..0a2854d 100644 --- a/dns/rdatatype.py +++ b/dns/rdatatype.py @@ -22,8 +22,10 @@ from typing import Dict import dns.enum import dns.exception + class RdataType(dns.enum.IntEnum): """DNS Rdata Type""" + TYPE0 = 0 NONE = 0 A = 1 @@ -122,13 +124,19 @@ class RdataType(dns.enum.IntEnum): def _unknown_exception_class(cls): return UnknownRdatatype + _registered_by_text: Dict[str, RdataType] = {} _registered_by_value: Dict[RdataType, str] = {} _metatypes = {RdataType.OPT} -_singletons = {RdataType.SOA, RdataType.NXT, RdataType.DNAME, - RdataType.NSEC, RdataType.CNAME} +_singletons = { + RdataType.SOA, + RdataType.NXT, + RdataType.DNAME, + RdataType.NSEC, + RdataType.CNAME, +} class UnknownRdatatype(dns.exception.DNSException): @@ -150,7 +158,7 @@ def from_text(text: str) -> RdataType: Returns a ``dns.rdatatype.RdataType``. """ - text = text.upper().replace('-', '_') + text = text.upper().replace("-", "_") try: return RdataType.from_text(text) except UnknownRdatatype: @@ -176,7 +184,7 @@ def to_text(value: RdataType) -> str: registered_text = _registered_by_value.get(value) if registered_text: text = registered_text - return text.replace('_', '-') + return text.replace("_", "-") def is_metatype(rdtype: RdataType) -> bool: @@ -211,8 +219,11 @@ def is_singleton(rdtype: RdataType) -> bool: return True return False + # pylint: disable=redefined-outer-name -def register_type(rdtype: RdataType, rdtype_text: str, is_singleton: bool=False) -> None: +def register_type( + rdtype: RdataType, rdtype_text: str, is_singleton: bool = False +) -> None: """Dynamically register an rdatatype. *rdtype*, a ``dns.rdatatype.RdataType``, the rdatatype to register. @@ -228,6 +239,7 @@ def register_type(rdtype: RdataType, rdtype_text: str, is_singleton: bool=False) if is_singleton: _singletons.add(rdtype) + ### BEGIN generated RdataType constants TYPE0 = RdataType.TYPE0 diff --git a/dns/rdtypes/ANY/AMTRELAY.py b/dns/rdtypes/ANY/AMTRELAY.py index 9f093de..dfe7abc 100644 --- a/dns/rdtypes/ANY/AMTRELAY.py +++ b/dns/rdtypes/ANY/AMTRELAY.py @@ -23,7 +23,7 @@ import dns.rdtypes.util class Relay(dns.rdtypes.util.Gateway): - name = 'AMTRELAY relay' + name = "AMTRELAY relay" @property def relay(self): @@ -37,10 +37,11 @@ class AMTRELAY(dns.rdata.Rdata): # see: RFC 8777 - __slots__ = ['precedence', 'discovery_optional', 'relay_type', 'relay'] + __slots__ = ["precedence", "discovery_optional", "relay_type", "relay"] - def __init__(self, rdclass, rdtype, precedence, discovery_optional, - relay_type, relay): + def __init__( + self, rdclass, rdtype, precedence, discovery_optional, relay_type, relay + ): super().__init__(rdclass, rdtype) relay = Relay(relay_type, relay) self.precedence = self._as_uint8(precedence) @@ -50,37 +51,42 @@ class AMTRELAY(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): relay = Relay(self.relay_type, self.relay).to_text(origin, relativize) - return '%d %d %d %s' % (self.precedence, self.discovery_optional, - self.relay_type, relay) + return "%d %d %d %s" % ( + self.precedence, + self.discovery_optional, + self.relay_type, + relay, + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): precedence = tok.get_uint8() discovery_optional = tok.get_uint8() if discovery_optional > 1: - raise dns.exception.SyntaxError('expecting 0 or 1') + raise dns.exception.SyntaxError("expecting 0 or 1") discovery_optional = bool(discovery_optional) relay_type = tok.get_uint8() - if relay_type > 0x7f: - raise dns.exception.SyntaxError('expecting an integer <= 127') - relay = Relay.from_text(relay_type, tok, origin, relativize, - relativize_to) - return cls(rdclass, rdtype, precedence, discovery_optional, relay_type, - relay.relay) + if relay_type > 0x7F: + raise dns.exception.SyntaxError("expecting an integer <= 127") + relay = Relay.from_text(relay_type, tok, origin, relativize, relativize_to) + return cls( + rdclass, rdtype, precedence, discovery_optional, relay_type, relay.relay + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): relay_type = self.relay_type | (self.discovery_optional << 7) header = struct.pack("!BB", self.precedence, relay_type) file.write(header) - Relay(self.relay_type, self.relay).to_wire(file, compress, origin, - canonicalize) + Relay(self.relay_type, self.relay).to_wire(file, compress, origin, canonicalize) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (precedence, relay_type) = parser.get_struct('!BB') + (precedence, relay_type) = parser.get_struct("!BB") discovery_optional = bool(relay_type >> 7) - relay_type &= 0x7f + relay_type &= 0x7F relay = Relay.from_wire_parser(relay_type, parser, origin) - return cls(rdclass, rdtype, precedence, discovery_optional, relay_type, - relay.relay) + return cls( + rdclass, rdtype, precedence, discovery_optional, relay_type, relay.relay + ) diff --git a/dns/rdtypes/ANY/CAA.py b/dns/rdtypes/ANY/CAA.py index c86b45e..8afb538 100644 --- a/dns/rdtypes/ANY/CAA.py +++ b/dns/rdtypes/ANY/CAA.py @@ -30,7 +30,7 @@ class CAA(dns.rdata.Rdata): # see: RFC 6844 - __slots__ = ['flags', 'tag', 'value'] + __slots__ = ["flags", "tag", "value"] def __init__(self, rdclass, rdtype, flags, tag, value): super().__init__(rdclass, rdtype) @@ -41,23 +41,26 @@ class CAA(dns.rdata.Rdata): self.value = self._as_bytes(value) def to_text(self, origin=None, relativize=True, **kw): - return '%u %s "%s"' % (self.flags, - dns.rdata._escapify(self.tag), - dns.rdata._escapify(self.value)) + return '%u %s "%s"' % ( + self.flags, + dns.rdata._escapify(self.tag), + dns.rdata._escapify(self.value), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): flags = tok.get_uint8() tag = tok.get_string().encode() value = tok.get_string().encode() return cls(rdclass, rdtype, flags, tag, value) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!B', self.flags)) + file.write(struct.pack("!B", self.flags)) l = len(self.tag) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.tag) file.write(self.value) diff --git a/dns/rdtypes/ANY/CDNSKEY.py b/dns/rdtypes/ANY/CDNSKEY.py index 7ea8f2a..869523f 100644 --- a/dns/rdtypes/ANY/CDNSKEY.py +++ b/dns/rdtypes/ANY/CDNSKEY.py @@ -19,9 +19,15 @@ import dns.rdtypes.dnskeybase # lgtm[py/import-and-import-from] import dns.immutable # pylint: disable=unused-import -from dns.rdtypes.dnskeybase import SEP, REVOKE, ZONE # noqa: F401 lgtm[py/unused-import] +from dns.rdtypes.dnskeybase import ( + SEP, + REVOKE, + ZONE, +) # noqa: F401 lgtm[py/unused-import] + # pylint: enable=unused-import + @dns.immutable.immutable class CDNSKEY(dns.rdtypes.dnskeybase.DNSKEYBase): diff --git a/dns/rdtypes/ANY/CERT.py b/dns/rdtypes/ANY/CERT.py index f8990eb..1b0cbec 100644 --- a/dns/rdtypes/ANY/CERT.py +++ b/dns/rdtypes/ANY/CERT.py @@ -25,29 +25,29 @@ import dns.rdata import dns.tokenizer _ctype_by_value = { - 1: 'PKIX', - 2: 'SPKI', - 3: 'PGP', - 4: 'IPKIX', - 5: 'ISPKI', - 6: 'IPGP', - 7: 'ACPKIX', - 8: 'IACPKIX', - 253: 'URI', - 254: 'OID', + 1: "PKIX", + 2: "SPKI", + 3: "PGP", + 4: "IPKIX", + 5: "ISPKI", + 6: "IPGP", + 7: "ACPKIX", + 8: "IACPKIX", + 253: "URI", + 254: "OID", } _ctype_by_name = { - 'PKIX': 1, - 'SPKI': 2, - 'PGP': 3, - 'IPKIX': 4, - 'ISPKI': 5, - 'IPGP': 6, - 'ACPKIX': 7, - 'IACPKIX': 8, - 'URI': 253, - 'OID': 254, + "PKIX": 1, + "SPKI": 2, + "PGP": 3, + "IPKIX": 4, + "ISPKI": 5, + "IPGP": 6, + "ACPKIX": 7, + "IACPKIX": 8, + "URI": 253, + "OID": 254, } @@ -72,10 +72,11 @@ class CERT(dns.rdata.Rdata): # see RFC 4398 - __slots__ = ['certificate_type', 'key_tag', 'algorithm', 'certificate'] + __slots__ = ["certificate_type", "key_tag", "algorithm", "certificate"] - def __init__(self, rdclass, rdtype, certificate_type, key_tag, algorithm, - certificate): + def __init__( + self, rdclass, rdtype, certificate_type, key_tag, algorithm, certificate + ): super().__init__(rdclass, rdtype) self.certificate_type = self._as_uint16(certificate_type) self.key_tag = self._as_uint16(key_tag) @@ -84,24 +85,28 @@ class CERT(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): certificate_type = _ctype_to_text(self.certificate_type) - return "%s %d %s %s" % (certificate_type, self.key_tag, - dns.dnssectypes.Algorithm.to_text(self.algorithm), - dns.rdata._base64ify(self.certificate, **kw)) + return "%s %d %s %s" % ( + certificate_type, + self.key_tag, + dns.dnssectypes.Algorithm.to_text(self.algorithm), + dns.rdata._base64ify(self.certificate, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): certificate_type = _ctype_from_text(tok.get_string()) key_tag = tok.get_uint16() algorithm = dns.dnssectypes.Algorithm.from_text(tok.get_string()) b64 = tok.concatenate_remaining_identifiers().encode() certificate = base64.b64decode(b64) - return cls(rdclass, rdtype, certificate_type, key_tag, - algorithm, certificate) + return cls(rdclass, rdtype, certificate_type, key_tag, algorithm, certificate) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - prefix = struct.pack("!HHB", self.certificate_type, self.key_tag, - self.algorithm) + prefix = struct.pack( + "!HHB", self.certificate_type, self.key_tag, self.algorithm + ) file.write(prefix) file.write(self.certificate) @@ -109,5 +114,4 @@ class CERT(dns.rdata.Rdata): def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): (certificate_type, key_tag, algorithm) = parser.get_struct("!HHB") certificate = parser.get_remaining() - return cls(rdclass, rdtype, certificate_type, key_tag, algorithm, - certificate) + return cls(rdclass, rdtype, certificate_type, key_tag, algorithm, certificate) diff --git a/dns/rdtypes/ANY/CSYNC.py b/dns/rdtypes/ANY/CSYNC.py index 979028a..f819c08 100644 --- a/dns/rdtypes/ANY/CSYNC.py +++ b/dns/rdtypes/ANY/CSYNC.py @@ -27,7 +27,7 @@ import dns.rdtypes.util @dns.immutable.immutable class Bitmap(dns.rdtypes.util.Bitmap): - type_name = 'CSYNC' + type_name = "CSYNC" @dns.immutable.immutable @@ -35,7 +35,7 @@ class CSYNC(dns.rdata.Rdata): """CSYNC record""" - __slots__ = ['serial', 'flags', 'windows'] + __slots__ = ["serial", "flags", "windows"] def __init__(self, rdclass, rdtype, serial, flags, windows): super().__init__(rdclass, rdtype) @@ -47,18 +47,19 @@ class CSYNC(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): text = Bitmap(self.windows).to_text() - return '%d %d%s' % (self.serial, self.flags, text) + return "%d %d%s" % (self.serial, self.flags, text) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): serial = tok.get_uint32() flags = tok.get_uint16() bitmap = Bitmap.from_text(tok) return cls(rdclass, rdtype, serial, flags, bitmap) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!IH', self.serial, self.flags)) + file.write(struct.pack("!IH", self.serial, self.flags)) Bitmap(self.windows).to_wire(file) @classmethod diff --git a/dns/rdtypes/ANY/DNSKEY.py b/dns/rdtypes/ANY/DNSKEY.py index cc0bf8c..50fa05b 100644 --- a/dns/rdtypes/ANY/DNSKEY.py +++ b/dns/rdtypes/ANY/DNSKEY.py @@ -19,9 +19,15 @@ import dns.rdtypes.dnskeybase # lgtm[py/import-and-import-from] import dns.immutable # pylint: disable=unused-import -from dns.rdtypes.dnskeybase import SEP, REVOKE, ZONE # noqa: F401 lgtm[py/unused-import] +from dns.rdtypes.dnskeybase import ( + SEP, + REVOKE, + ZONE, +) # noqa: F401 lgtm[py/unused-import] + # pylint: enable=unused-import + @dns.immutable.immutable class DNSKEY(dns.rdtypes.dnskeybase.DNSKEYBase): diff --git a/dns/rdtypes/ANY/GPOS.py b/dns/rdtypes/ANY/GPOS.py index 29fa8f8..30aab32 100644 --- a/dns/rdtypes/ANY/GPOS.py +++ b/dns/rdtypes/ANY/GPOS.py @@ -26,19 +26,19 @@ import dns.tokenizer def _validate_float_string(what): if len(what) == 0: raise dns.exception.FormError - if what[0] == b'-'[0] or what[0] == b'+'[0]: + if what[0] == b"-"[0] or what[0] == b"+"[0]: what = what[1:] if what.isdigit(): return try: - (left, right) = what.split(b'.') + (left, right) = what.split(b".") except ValueError: raise dns.exception.FormError - if left == b'' and right == b'': + if left == b"" and right == b"": raise dns.exception.FormError - if not left == b'' and not left.decode().isdigit(): + if not left == b"" and not left.decode().isdigit(): raise dns.exception.FormError - if not right == b'' and not right.decode().isdigit(): + if not right == b"" and not right.decode().isdigit(): raise dns.exception.FormError @@ -49,18 +49,15 @@ class GPOS(dns.rdata.Rdata): # see: RFC 1712 - __slots__ = ['latitude', 'longitude', 'altitude'] + __slots__ = ["latitude", "longitude", "altitude"] def __init__(self, rdclass, rdtype, latitude, longitude, altitude): super().__init__(rdclass, rdtype) - if isinstance(latitude, float) or \ - isinstance(latitude, int): + if isinstance(latitude, float) or isinstance(latitude, int): latitude = str(latitude) - if isinstance(longitude, float) or \ - isinstance(longitude, int): + if isinstance(longitude, float) or isinstance(longitude, int): longitude = str(longitude) - if isinstance(altitude, float) or \ - isinstance(altitude, int): + if isinstance(altitude, float) or isinstance(altitude, int): altitude = str(altitude) latitude = self._as_bytes(latitude, True, 255) longitude = self._as_bytes(longitude, True, 255) @@ -73,19 +70,20 @@ class GPOS(dns.rdata.Rdata): self.altitude = altitude flat = self.float_latitude if flat < -90.0 or flat > 90.0: - raise dns.exception.FormError('bad latitude') + raise dns.exception.FormError("bad latitude") flong = self.float_longitude if flong < -180.0 or flong > 180.0: - raise dns.exception.FormError('bad longitude') + raise dns.exception.FormError("bad longitude") def to_text(self, origin=None, relativize=True, **kw): - return '{} {} {}'.format(self.latitude.decode(), - self.longitude.decode(), - self.altitude.decode()) + return "{} {} {}".format( + self.latitude.decode(), self.longitude.decode(), self.altitude.decode() + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): latitude = tok.get_string() longitude = tok.get_string() altitude = tok.get_string() @@ -94,15 +92,15 @@ class GPOS(dns.rdata.Rdata): def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.latitude) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.latitude) l = len(self.longitude) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.longitude) l = len(self.altitude) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.altitude) @classmethod diff --git a/dns/rdtypes/ANY/HINFO.py b/dns/rdtypes/ANY/HINFO.py index cd04969..513c155 100644 --- a/dns/rdtypes/ANY/HINFO.py +++ b/dns/rdtypes/ANY/HINFO.py @@ -30,7 +30,7 @@ class HINFO(dns.rdata.Rdata): # see: RFC 1035 - __slots__ = ['cpu', 'os'] + __slots__ = ["cpu", "os"] def __init__(self, rdclass, rdtype, cpu, os): super().__init__(rdclass, rdtype) @@ -38,12 +38,14 @@ class HINFO(dns.rdata.Rdata): self.os = self._as_bytes(os, True, 255) def to_text(self, origin=None, relativize=True, **kw): - return '"{}" "{}"'.format(dns.rdata._escapify(self.cpu), - dns.rdata._escapify(self.os)) + return '"{}" "{}"'.format( + dns.rdata._escapify(self.cpu), dns.rdata._escapify(self.os) + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): cpu = tok.get_string(max_length=255) os = tok.get_string(max_length=255) return cls(rdclass, rdtype, cpu, os) @@ -51,11 +53,11 @@ class HINFO(dns.rdata.Rdata): def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.cpu) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.cpu) l = len(self.os) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.os) @classmethod diff --git a/dns/rdtypes/ANY/HIP.py b/dns/rdtypes/ANY/HIP.py index e887359..01fec82 100644 --- a/dns/rdtypes/ANY/HIP.py +++ b/dns/rdtypes/ANY/HIP.py @@ -32,7 +32,7 @@ class HIP(dns.rdata.Rdata): # see: RFC 5205 - __slots__ = ['hit', 'algorithm', 'key', 'servers'] + __slots__ = ["hit", "algorithm", "key", "servers"] def __init__(self, rdclass, rdtype, hit, algorithm, key, servers): super().__init__(rdclass, rdtype) @@ -43,18 +43,19 @@ class HIP(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): hit = binascii.hexlify(self.hit).decode() - key = base64.b64encode(self.key).replace(b'\n', b'').decode() - text = '' + key = base64.b64encode(self.key).replace(b"\n", b"").decode() + text = "" servers = [] for server in self.servers: servers.append(server.choose_relativity(origin, relativize)) if len(servers) > 0: - text += (' ' + ' '.join((x.to_unicode() for x in servers))) - return '%u %s %s%s' % (self.algorithm, hit, key, text) + text += " " + " ".join((x.to_unicode() for x in servers)) + return "%u %s %s%s" % (self.algorithm, hit, key, text) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_uint8() hit = binascii.unhexlify(tok.get_string().encode()) key = base64.b64decode(tok.get_string().encode()) @@ -75,7 +76,7 @@ class HIP(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (lh, algorithm, lk) = parser.get_struct('!BBH') + (lh, algorithm, lk) = parser.get_struct("!BBH") hit = parser.get_bytes(lh) key = parser.get_bytes(lk) servers = [] diff --git a/dns/rdtypes/ANY/ISDN.py b/dns/rdtypes/ANY/ISDN.py index b9a49ad..536a35d 100644 --- a/dns/rdtypes/ANY/ISDN.py +++ b/dns/rdtypes/ANY/ISDN.py @@ -30,7 +30,7 @@ class ISDN(dns.rdata.Rdata): # see: RFC 1183 - __slots__ = ['address', 'subaddress'] + __slots__ = ["address", "subaddress"] def __init__(self, rdclass, rdtype, address, subaddress): super().__init__(rdclass, rdtype) @@ -39,31 +39,33 @@ class ISDN(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): if self.subaddress: - return '"{}" "{}"'.format(dns.rdata._escapify(self.address), - dns.rdata._escapify(self.subaddress)) + return '"{}" "{}"'.format( + dns.rdata._escapify(self.address), dns.rdata._escapify(self.subaddress) + ) else: return '"%s"' % dns.rdata._escapify(self.address) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_string() tokens = tok.get_remaining(max_tokens=1) if len(tokens) >= 1: subaddress = tokens[0].unescape().value else: - subaddress = '' + subaddress = "" return cls(rdclass, rdtype, address, subaddress) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.address) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.address) l = len(self.subaddress) if l > 0: assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.subaddress) @classmethod @@ -72,5 +74,5 @@ class ISDN(dns.rdata.Rdata): if parser.remaining() > 0: subaddress = parser.get_counted_bytes() else: - subaddress = b'' + subaddress = b"" return cls(rdclass, rdtype, address, subaddress) diff --git a/dns/rdtypes/ANY/L32.py b/dns/rdtypes/ANY/L32.py index 038fc3a..14be01f 100644 --- a/dns/rdtypes/ANY/L32.py +++ b/dns/rdtypes/ANY/L32.py @@ -13,7 +13,7 @@ class L32(dns.rdata.Rdata): # see: rfc6742.txt - __slots__ = ['preference', 'locator32'] + __slots__ = ["preference", "locator32"] def __init__(self, rdclass, rdtype, preference, locator32): super().__init__(rdclass, rdtype) @@ -21,17 +21,18 @@ class L32(dns.rdata.Rdata): self.locator32 = self._as_ipv4_address(locator32) def to_text(self, origin=None, relativize=True, **kw): - return f'{self.preference} {self.locator32}' + return f"{self.preference} {self.locator32}" @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() nodeid = tok.get_identifier() return cls(rdclass, rdtype, preference, nodeid) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!H', self.preference)) + file.write(struct.pack("!H", self.preference)) file.write(dns.ipv4.inet_aton(self.locator32)) @classmethod diff --git a/dns/rdtypes/ANY/L64.py b/dns/rdtypes/ANY/L64.py index aab36a8..d083d40 100644 --- a/dns/rdtypes/ANY/L64.py +++ b/dns/rdtypes/ANY/L64.py @@ -13,33 +13,33 @@ class L64(dns.rdata.Rdata): # see: rfc6742.txt - __slots__ = ['preference', 'locator64'] + __slots__ = ["preference", "locator64"] def __init__(self, rdclass, rdtype, preference, locator64): super().__init__(rdclass, rdtype) self.preference = self._as_uint16(preference) if isinstance(locator64, bytes): if len(locator64) != 8: - raise ValueError('invalid locator64') - self.locator64 = dns.rdata._hexify(locator64, 4, b':') + raise ValueError("invalid locator64") + self.locator64 = dns.rdata._hexify(locator64, 4, b":") else: - dns.rdtypes.util.parse_formatted_hex(locator64, 4, 4, ':') + dns.rdtypes.util.parse_formatted_hex(locator64, 4, 4, ":") self.locator64 = locator64 def to_text(self, origin=None, relativize=True, **kw): - return f'{self.preference} {self.locator64}' + return f"{self.preference} {self.locator64}" @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() locator64 = tok.get_identifier() return cls(rdclass, rdtype, preference, locator64) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!H', self.preference)) - file.write(dns.rdtypes.util.parse_formatted_hex(self.locator64, - 4, 4, ':')) + file.write(struct.pack("!H", self.preference)) + file.write(dns.rdtypes.util.parse_formatted_hex(self.locator64, 4, 4, ":")) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): diff --git a/dns/rdtypes/ANY/LOC.py b/dns/rdtypes/ANY/LOC.py index c939899..52c9753 100644 --- a/dns/rdtypes/ANY/LOC.py +++ b/dns/rdtypes/ANY/LOC.py @@ -93,15 +93,15 @@ def _decode_size(what, desc): def _check_coordinate_list(value, low, high): if value[0] < low or value[0] > high: - raise ValueError(f'not in range [{low}, {high}]') + raise ValueError(f"not in range [{low}, {high}]") if value[1] < 0 or value[1] > 59: - raise ValueError('bad minutes value') + raise ValueError("bad minutes value") if value[2] < 0 or value[2] > 59: - raise ValueError('bad seconds value') + raise ValueError("bad seconds value") if value[3] < 0 or value[3] > 999: - raise ValueError('bad milliseconds value') + raise ValueError("bad milliseconds value") if value[4] != 1 and value[4] != -1: - raise ValueError('bad hemisphere value') + raise ValueError("bad hemisphere value") @dns.immutable.immutable @@ -111,12 +111,26 @@ class LOC(dns.rdata.Rdata): # see: RFC 1876 - __slots__ = ['latitude', 'longitude', 'altitude', 'size', - 'horizontal_precision', 'vertical_precision'] - - def __init__(self, rdclass, rdtype, latitude, longitude, altitude, - size=_default_size, hprec=_default_hprec, - vprec=_default_vprec): + __slots__ = [ + "latitude", + "longitude", + "altitude", + "size", + "horizontal_precision", + "vertical_precision", + ] + + def __init__( + self, + rdclass, + rdtype, + latitude, + longitude, + altitude, + size=_default_size, + hprec=_default_hprec, + vprec=_default_vprec, + ): """Initialize a LOC record instance. The parameters I{latitude} and I{longitude} may be either a 4-tuple @@ -145,34 +159,44 @@ class LOC(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): if self.latitude[4] > 0: - lat_hemisphere = 'N' + lat_hemisphere = "N" else: - lat_hemisphere = 'S' + lat_hemisphere = "S" if self.longitude[4] > 0: - long_hemisphere = 'E' + long_hemisphere = "E" else: - long_hemisphere = 'W' + long_hemisphere = "W" text = "%d %d %d.%03d %s %d %d %d.%03d %s %0.2fm" % ( - self.latitude[0], self.latitude[1], - self.latitude[2], self.latitude[3], lat_hemisphere, - self.longitude[0], self.longitude[1], self.longitude[2], - self.longitude[3], long_hemisphere, - self.altitude / 100.0 + self.latitude[0], + self.latitude[1], + self.latitude[2], + self.latitude[3], + lat_hemisphere, + self.longitude[0], + self.longitude[1], + self.longitude[2], + self.longitude[3], + long_hemisphere, + self.altitude / 100.0, ) # do not print default values - if self.size != _default_size or \ - self.horizontal_precision != _default_hprec or \ - self.vertical_precision != _default_vprec: + if ( + self.size != _default_size + or self.horizontal_precision != _default_hprec + or self.vertical_precision != _default_vprec + ): text += " {:0.2f}m {:0.2f}m {:0.2f}m".format( - self.size / 100.0, self.horizontal_precision / 100.0, - self.vertical_precision / 100.0 + self.size / 100.0, + self.horizontal_precision / 100.0, + self.vertical_precision / 100.0, ) return text @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): latitude = [0, 0, 0, 0, 1] longitude = [0, 0, 0, 0, 1] size = _default_size @@ -184,16 +208,14 @@ class LOC(dns.rdata.Rdata): if t.isdigit(): latitude[1] = int(t) t = tok.get_string() - if '.' in t: - (seconds, milliseconds) = t.split('.') + if "." in t: + (seconds, milliseconds) = t.split(".") if not seconds.isdigit(): - raise dns.exception.SyntaxError( - 'bad latitude seconds value') + raise dns.exception.SyntaxError("bad latitude seconds value") latitude[2] = int(seconds) l = len(milliseconds) if l == 0 or l > 3 or not milliseconds.isdigit(): - raise dns.exception.SyntaxError( - 'bad latitude milliseconds value') + raise dns.exception.SyntaxError("bad latitude milliseconds value") if l == 1: m = 100 elif l == 2: @@ -205,26 +227,24 @@ class LOC(dns.rdata.Rdata): elif t.isdigit(): latitude[2] = int(t) t = tok.get_string() - if t == 'S': + if t == "S": latitude[4] = -1 - elif t != 'N': - raise dns.exception.SyntaxError('bad latitude hemisphere value') + elif t != "N": + raise dns.exception.SyntaxError("bad latitude hemisphere value") longitude[0] = tok.get_int() t = tok.get_string() if t.isdigit(): longitude[1] = int(t) t = tok.get_string() - if '.' in t: - (seconds, milliseconds) = t.split('.') + if "." in t: + (seconds, milliseconds) = t.split(".") if not seconds.isdigit(): - raise dns.exception.SyntaxError( - 'bad longitude seconds value') + raise dns.exception.SyntaxError("bad longitude seconds value") longitude[2] = int(seconds) l = len(milliseconds) if l == 0 or l > 3 or not milliseconds.isdigit(): - raise dns.exception.SyntaxError( - 'bad longitude milliseconds value') + raise dns.exception.SyntaxError("bad longitude milliseconds value") if l == 1: m = 100 elif l == 2: @@ -236,64 +256,75 @@ class LOC(dns.rdata.Rdata): elif t.isdigit(): longitude[2] = int(t) t = tok.get_string() - if t == 'W': + if t == "W": longitude[4] = -1 - elif t != 'E': - raise dns.exception.SyntaxError('bad longitude hemisphere value') + elif t != "E": + raise dns.exception.SyntaxError("bad longitude hemisphere value") t = tok.get_string() - if t[-1] == 'm': - t = t[0: -1] - altitude = float(t) * 100.0 # m -> cm + if t[-1] == "m": + t = t[0:-1] + altitude = float(t) * 100.0 # m -> cm tokens = tok.get_remaining(max_tokens=3) if len(tokens) >= 1: value = tokens[0].unescape().value - if value[-1] == 'm': - value = value[0: -1] - size = float(value) * 100.0 # m -> cm + if value[-1] == "m": + value = value[0:-1] + size = float(value) * 100.0 # m -> cm if len(tokens) >= 2: value = tokens[1].unescape().value - if value[-1] == 'm': - value = value[0: -1] - hprec = float(value) * 100.0 # m -> cm + if value[-1] == "m": + value = value[0:-1] + hprec = float(value) * 100.0 # m -> cm if len(tokens) >= 3: value = tokens[2].unescape().value - if value[-1] == 'm': - value = value[0: -1] - vprec = float(value) * 100.0 # m -> cm + if value[-1] == "m": + value = value[0:-1] + vprec = float(value) * 100.0 # m -> cm # Try encoding these now so we raise if they are bad _encode_size(size, "size") _encode_size(hprec, "horizontal precision") _encode_size(vprec, "vertical precision") - return cls(rdclass, rdtype, latitude, longitude, altitude, - size, hprec, vprec) + return cls(rdclass, rdtype, latitude, longitude, altitude, size, hprec, vprec) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - milliseconds = (self.latitude[0] * 3600000 + - self.latitude[1] * 60000 + - self.latitude[2] * 1000 + - self.latitude[3]) * self.latitude[4] + milliseconds = ( + self.latitude[0] * 3600000 + + self.latitude[1] * 60000 + + self.latitude[2] * 1000 + + self.latitude[3] + ) * self.latitude[4] latitude = 0x80000000 + milliseconds - milliseconds = (self.longitude[0] * 3600000 + - self.longitude[1] * 60000 + - self.longitude[2] * 1000 + - self.longitude[3]) * self.longitude[4] + milliseconds = ( + self.longitude[0] * 3600000 + + self.longitude[1] * 60000 + + self.longitude[2] * 1000 + + self.longitude[3] + ) * self.longitude[4] longitude = 0x80000000 + milliseconds altitude = int(self.altitude) + 10000000 size = _encode_size(self.size, "size") hprec = _encode_size(self.horizontal_precision, "horizontal precision") vprec = _encode_size(self.vertical_precision, "vertical precision") - wire = struct.pack("!BBBBIII", 0, size, hprec, vprec, latitude, - longitude, altitude) + wire = struct.pack( + "!BBBBIII", 0, size, hprec, vprec, latitude, longitude, altitude + ) file.write(wire) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (version, size, hprec, vprec, latitude, longitude, altitude) = \ - parser.get_struct("!BBBBIII") + ( + version, + size, + hprec, + vprec, + latitude, + longitude, + altitude, + ) = parser.get_struct("!BBBBIII") if version != 0: raise dns.exception.FormError("LOC version not zero") if latitude < _MIN_LATITUDE or latitude > _MAX_LATITUDE: @@ -312,8 +343,7 @@ class LOC(dns.rdata.Rdata): size = _decode_size(size, "size") hprec = _decode_size(hprec, "horizontal precision") vprec = _decode_size(vprec, "vertical precision") - return cls(rdclass, rdtype, latitude, longitude, altitude, - size, hprec, vprec) + return cls(rdclass, rdtype, latitude, longitude, altitude, size, hprec, vprec) @property def float_latitude(self): diff --git a/dns/rdtypes/ANY/LP.py b/dns/rdtypes/ANY/LP.py index a4adffb..8a7c512 100644 --- a/dns/rdtypes/ANY/LP.py +++ b/dns/rdtypes/ANY/LP.py @@ -13,7 +13,7 @@ class LP(dns.rdata.Rdata): # see: rfc6742.txt - __slots__ = ['preference', 'fqdn'] + __slots__ = ["preference", "fqdn"] def __init__(self, rdclass, rdtype, preference, fqdn): super().__init__(rdclass, rdtype) @@ -22,17 +22,18 @@ class LP(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): fqdn = self.fqdn.choose_relativity(origin, relativize) - return '%d %s' % (self.preference, fqdn) + return "%d %s" % (self.preference, fqdn) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() fqdn = tok.get_name(origin, relativize, relativize_to) return cls(rdclass, rdtype, preference, fqdn) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!H', self.preference)) + file.write(struct.pack("!H", self.preference)) self.fqdn.to_wire(file, compress, origin, canonicalize) @classmethod diff --git a/dns/rdtypes/ANY/NID.py b/dns/rdtypes/ANY/NID.py index 74951bb..ad54aca 100644 --- a/dns/rdtypes/ANY/NID.py +++ b/dns/rdtypes/ANY/NID.py @@ -13,32 +13,33 @@ class NID(dns.rdata.Rdata): # see: rfc6742.txt - __slots__ = ['preference', 'nodeid'] + __slots__ = ["preference", "nodeid"] def __init__(self, rdclass, rdtype, preference, nodeid): super().__init__(rdclass, rdtype) self.preference = self._as_uint16(preference) if isinstance(nodeid, bytes): if len(nodeid) != 8: - raise ValueError('invalid nodeid') - self.nodeid = dns.rdata._hexify(nodeid, 4, b':') + raise ValueError("invalid nodeid") + self.nodeid = dns.rdata._hexify(nodeid, 4, b":") else: - dns.rdtypes.util.parse_formatted_hex(nodeid, 4, 4, ':') + dns.rdtypes.util.parse_formatted_hex(nodeid, 4, 4, ":") self.nodeid = nodeid def to_text(self, origin=None, relativize=True, **kw): - return f'{self.preference} {self.nodeid}' + return f"{self.preference} {self.nodeid}" @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() nodeid = tok.get_identifier() return cls(rdclass, rdtype, preference, nodeid) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - file.write(struct.pack('!H', self.preference)) - file.write(dns.rdtypes.util.parse_formatted_hex(self.nodeid, 4, 4, ':')) + file.write(struct.pack("!H", self.preference)) + file.write(dns.rdtypes.util.parse_formatted_hex(self.nodeid, 4, 4, ":")) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): diff --git a/dns/rdtypes/ANY/NSEC.py b/dns/rdtypes/ANY/NSEC.py index dc31f4c..7af7b77 100644 --- a/dns/rdtypes/ANY/NSEC.py +++ b/dns/rdtypes/ANY/NSEC.py @@ -25,7 +25,7 @@ import dns.rdtypes.util @dns.immutable.immutable class Bitmap(dns.rdtypes.util.Bitmap): - type_name = 'NSEC' + type_name = "NSEC" @dns.immutable.immutable @@ -33,7 +33,7 @@ class NSEC(dns.rdata.Rdata): """NSEC record""" - __slots__ = ['next', 'windows'] + __slots__ = ["next", "windows"] def __init__(self, rdclass, rdtype, next, windows): super().__init__(rdclass, rdtype) @@ -45,11 +45,12 @@ class NSEC(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): next = self.next.choose_relativity(origin, relativize) text = Bitmap(self.windows).to_text() - return '{}{}'.format(next, text) + return "{}{}".format(next, text) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): next = tok.get_name(origin, relativize, relativize_to) windows = Bitmap.from_text(tok) return cls(rdclass, rdtype, next, windows) diff --git a/dns/rdtypes/ANY/NSEC3.py b/dns/rdtypes/ANY/NSEC3.py index 14242bd..6eae16e 100644 --- a/dns/rdtypes/ANY/NSEC3.py +++ b/dns/rdtypes/ANY/NSEC3.py @@ -26,10 +26,12 @@ import dns.rdatatype import dns.rdtypes.util -b32_hex_to_normal = bytes.maketrans(b'0123456789ABCDEFGHIJKLMNOPQRSTUV', - b'ABCDEFGHIJKLMNOPQRSTUVWXYZ234567') -b32_normal_to_hex = bytes.maketrans(b'ABCDEFGHIJKLMNOPQRSTUVWXYZ234567', - b'0123456789ABCDEFGHIJKLMNOPQRSTUV') +b32_hex_to_normal = bytes.maketrans( + b"0123456789ABCDEFGHIJKLMNOPQRSTUV", b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567" +) +b32_normal_to_hex = bytes.maketrans( + b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567", b"0123456789ABCDEFGHIJKLMNOPQRSTUV" +) # hash algorithm constants SHA1 = 1 @@ -40,7 +42,7 @@ OPTOUT = 1 @dns.immutable.immutable class Bitmap(dns.rdtypes.util.Bitmap): - type_name = 'NSEC3' + type_name = "NSEC3" @dns.immutable.immutable @@ -48,10 +50,11 @@ class NSEC3(dns.rdata.Rdata): """NSEC3 record""" - __slots__ = ['algorithm', 'flags', 'iterations', 'salt', 'next', 'windows'] + __slots__ = ["algorithm", "flags", "iterations", "salt", "next", "windows"] - def __init__(self, rdclass, rdtype, algorithm, flags, iterations, salt, - next, windows): + def __init__( + self, rdclass, rdtype, algorithm, flags, iterations, salt, next, windows + ): super().__init__(rdclass, rdtype) self.algorithm = self._as_uint8(algorithm) self.flags = self._as_uint8(flags) @@ -63,38 +66,41 @@ class NSEC3(dns.rdata.Rdata): self.windows = tuple(windows.windows) def to_text(self, origin=None, relativize=True, **kw): - next = base64.b32encode(self.next).translate( - b32_normal_to_hex).lower().decode() - if self.salt == b'': - salt = '-' + next = base64.b32encode(self.next).translate(b32_normal_to_hex).lower().decode() + if self.salt == b"": + salt = "-" else: salt = binascii.hexlify(self.salt).decode() text = Bitmap(self.windows).to_text() - return '%u %u %u %s %s%s' % (self.algorithm, self.flags, - self.iterations, salt, next, text) + return "%u %u %u %s %s%s" % ( + self.algorithm, + self.flags, + self.iterations, + salt, + next, + text, + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_uint8() flags = tok.get_uint8() iterations = tok.get_uint16() salt = tok.get_string() - if salt == '-': - salt = b'' + if salt == "-": + salt = b"" else: - salt = binascii.unhexlify(salt.encode('ascii')) - next = tok.get_string().encode( - 'ascii').upper().translate(b32_hex_to_normal) + salt = binascii.unhexlify(salt.encode("ascii")) + next = tok.get_string().encode("ascii").upper().translate(b32_hex_to_normal) next = base64.b32decode(next) bitmap = Bitmap.from_text(tok) - return cls(rdclass, rdtype, algorithm, flags, iterations, salt, next, - bitmap) + return cls(rdclass, rdtype, algorithm, flags, iterations, salt, next, bitmap) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.salt) - file.write(struct.pack("!BBHB", self.algorithm, self.flags, - self.iterations, l)) + file.write(struct.pack("!BBHB", self.algorithm, self.flags, self.iterations, l)) file.write(self.salt) l = len(self.next) file.write(struct.pack("!B", l)) @@ -103,9 +109,8 @@ class NSEC3(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (algorithm, flags, iterations) = parser.get_struct('!BBH') + (algorithm, flags, iterations) = parser.get_struct("!BBH") salt = parser.get_counted_bytes() next = parser.get_counted_bytes() bitmap = Bitmap.from_wire_parser(parser) - return cls(rdclass, rdtype, algorithm, flags, iterations, salt, next, - bitmap) + return cls(rdclass, rdtype, algorithm, flags, iterations, salt, next, bitmap) diff --git a/dns/rdtypes/ANY/NSEC3PARAM.py b/dns/rdtypes/ANY/NSEC3PARAM.py index 299bf6e..1b7269a 100644 --- a/dns/rdtypes/ANY/NSEC3PARAM.py +++ b/dns/rdtypes/ANY/NSEC3PARAM.py @@ -28,7 +28,7 @@ class NSEC3PARAM(dns.rdata.Rdata): """NSEC3PARAM record""" - __slots__ = ['algorithm', 'flags', 'iterations', 'salt'] + __slots__ = ["algorithm", "flags", "iterations", "salt"] def __init__(self, rdclass, rdtype, algorithm, flags, iterations, salt): super().__init__(rdclass, rdtype) @@ -38,34 +38,33 @@ class NSEC3PARAM(dns.rdata.Rdata): self.salt = self._as_bytes(salt, True, 255) def to_text(self, origin=None, relativize=True, **kw): - if self.salt == b'': - salt = '-' + if self.salt == b"": + salt = "-" else: salt = binascii.hexlify(self.salt).decode() - return '%u %u %u %s' % (self.algorithm, self.flags, self.iterations, - salt) + return "%u %u %u %s" % (self.algorithm, self.flags, self.iterations, salt) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_uint8() flags = tok.get_uint8() iterations = tok.get_uint16() salt = tok.get_string() - if salt == '-': - salt = '' + if salt == "-": + salt = "" else: salt = binascii.unhexlify(salt.encode()) return cls(rdclass, rdtype, algorithm, flags, iterations, salt) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.salt) - file.write(struct.pack("!BBHB", self.algorithm, self.flags, - self.iterations, l)) + file.write(struct.pack("!BBHB", self.algorithm, self.flags, self.iterations, l)) file.write(self.salt) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (algorithm, flags, iterations) = parser.get_struct('!BBH') + (algorithm, flags, iterations) = parser.get_struct("!BBH") salt = parser.get_counted_bytes() return cls(rdclass, rdtype, algorithm, flags, iterations, salt) diff --git a/dns/rdtypes/ANY/OPENPGPKEY.py b/dns/rdtypes/ANY/OPENPGPKEY.py index dcfa028..e5e2572 100644 --- a/dns/rdtypes/ANY/OPENPGPKEY.py +++ b/dns/rdtypes/ANY/OPENPGPKEY.py @@ -22,6 +22,7 @@ import dns.immutable import dns.rdata import dns.tokenizer + @dns.immutable.immutable class OPENPGPKEY(dns.rdata.Rdata): @@ -37,8 +38,9 @@ class OPENPGPKEY(dns.rdata.Rdata): return dns.rdata._base64ify(self.key, chunksize=None, **kw) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): b64 = tok.concatenate_remaining_identifiers().encode() key = base64.b64decode(b64) return cls(rdclass, rdtype, key) diff --git a/dns/rdtypes/ANY/OPT.py b/dns/rdtypes/ANY/OPT.py index 69b8fe7..36d4c7c 100644 --- a/dns/rdtypes/ANY/OPT.py +++ b/dns/rdtypes/ANY/OPT.py @@ -26,12 +26,13 @@ import dns.rdata # We don't implement from_text, and that's ok. # pylint: disable=abstract-method + @dns.immutable.immutable class OPT(dns.rdata.Rdata): """OPT record""" - __slots__ = ['options'] + __slots__ = ["options"] def __init__(self, rdclass, rdtype, options): """Initialize an OPT rdata. @@ -45,10 +46,12 @@ class OPT(dns.rdata.Rdata): """ super().__init__(rdclass, rdtype) + def as_option(option): if not isinstance(option, dns.edns.Option): - raise ValueError('option is not a dns.edns.option') + raise ValueError("option is not a dns.edns.option") return option + self.options = self._as_tuple(options, as_option) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): @@ -58,13 +61,13 @@ class OPT(dns.rdata.Rdata): file.write(owire) def to_text(self, origin=None, relativize=True, **kw): - return ' '.join(opt.to_text() for opt in self.options) + return " ".join(opt.to_text() for opt in self.options) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): options = [] while parser.remaining() > 0: - (otype, olen) = parser.get_struct('!HH') + (otype, olen) = parser.get_struct("!HH") with parser.restrict_to(olen): opt = dns.edns.option_from_wire_parser(otype, parser) options.append(opt) diff --git a/dns/rdtypes/ANY/RP.py b/dns/rdtypes/ANY/RP.py index a4e2297..c0c316b 100644 --- a/dns/rdtypes/ANY/RP.py +++ b/dns/rdtypes/ANY/RP.py @@ -28,7 +28,7 @@ class RP(dns.rdata.Rdata): # see: RFC 1183 - __slots__ = ['mbox', 'txt'] + __slots__ = ["mbox", "txt"] def __init__(self, rdclass, rdtype, mbox, txt): super().__init__(rdclass, rdtype) @@ -41,8 +41,9 @@ class RP(dns.rdata.Rdata): return "{} {}".format(str(mbox), str(txt)) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): mbox = tok.get_name(origin, relativize, relativize_to) txt = tok.get_name(origin, relativize, relativize_to) return cls(rdclass, rdtype, mbox, txt) diff --git a/dns/rdtypes/ANY/RRSIG.py b/dns/rdtypes/ANY/RRSIG.py index 82650c0..3d5ad0f 100644 --- a/dns/rdtypes/ANY/RRSIG.py +++ b/dns/rdtypes/ANY/RRSIG.py @@ -43,12 +43,11 @@ def sigtime_to_posixtime(what): hour = int(what[8:10]) minute = int(what[10:12]) second = int(what[12:14]) - return calendar.timegm((year, month, day, hour, minute, second, - 0, 0, 0)) + return calendar.timegm((year, month, day, hour, minute, second, 0, 0, 0)) def posixtime_to_sigtime(what): - return time.strftime('%Y%m%d%H%M%S', time.gmtime(what)) + return time.strftime("%Y%m%d%H%M%S", time.gmtime(what)) @dns.immutable.immutable @@ -56,13 +55,32 @@ class RRSIG(dns.rdata.Rdata): """RRSIG record""" - __slots__ = ['type_covered', 'algorithm', 'labels', 'original_ttl', - 'expiration', 'inception', 'key_tag', 'signer', - 'signature'] - - def __init__(self, rdclass, rdtype, type_covered, algorithm, labels, - original_ttl, expiration, inception, key_tag, signer, - signature): + __slots__ = [ + "type_covered", + "algorithm", + "labels", + "original_ttl", + "expiration", + "inception", + "key_tag", + "signer", + "signature", + ] + + def __init__( + self, + rdclass, + rdtype, + type_covered, + algorithm, + labels, + original_ttl, + expiration, + inception, + key_tag, + signer, + signature, + ): super().__init__(rdclass, rdtype) self.type_covered = self._as_rdatatype(type_covered) self.algorithm = dns.dnssectypes.Algorithm.make(algorithm) @@ -78,7 +96,7 @@ class RRSIG(dns.rdata.Rdata): return self.type_covered def to_text(self, origin=None, relativize=True, **kw): - return '%s %d %d %d %s %s %d %s %s' % ( + return "%s %d %d %d %s %s %d %s %s" % ( dns.rdatatype.to_text(self.type_covered), self.algorithm, self.labels, @@ -87,12 +105,13 @@ class RRSIG(dns.rdata.Rdata): posixtime_to_sigtime(self.inception), self.key_tag, self.signer.choose_relativity(origin, relativize), - dns.rdata._base64ify(self.signature, **kw) + dns.rdata._base64ify(self.signature, **kw), ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): type_covered = dns.rdatatype.from_text(tok.get_string()) algorithm = dns.dnssectypes.Algorithm.from_text(tok.get_string()) labels = tok.get_int() @@ -103,22 +122,38 @@ class RRSIG(dns.rdata.Rdata): signer = tok.get_name(origin, relativize, relativize_to) b64 = tok.concatenate_remaining_identifiers().encode() signature = base64.b64decode(b64) - return cls(rdclass, rdtype, type_covered, algorithm, labels, - original_ttl, expiration, inception, key_tag, signer, - signature) + return cls( + rdclass, + rdtype, + type_covered, + algorithm, + labels, + original_ttl, + expiration, + inception, + key_tag, + signer, + signature, + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - header = struct.pack('!HBBIIIH', self.type_covered, - self.algorithm, self.labels, - self.original_ttl, self.expiration, - self.inception, self.key_tag) + header = struct.pack( + "!HBBIIIH", + self.type_covered, + self.algorithm, + self.labels, + self.original_ttl, + self.expiration, + self.inception, + self.key_tag, + ) file.write(header) self.signer.to_wire(file, None, origin, canonicalize) file.write(self.signature) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - header = parser.get_struct('!HBBIIIH') + header = parser.get_struct("!HBBIIIH") signer = parser.get_name(origin) signature = parser.get_remaining() return cls(rdclass, rdtype, *header, signer, signature) diff --git a/dns/rdtypes/ANY/SOA.py b/dns/rdtypes/ANY/SOA.py index 7ce8865..6f6fe58 100644 --- a/dns/rdtypes/ANY/SOA.py +++ b/dns/rdtypes/ANY/SOA.py @@ -30,11 +30,11 @@ class SOA(dns.rdata.Rdata): # see: RFC 1035 - __slots__ = ['mname', 'rname', 'serial', 'refresh', 'retry', 'expire', - 'minimum'] + __slots__ = ["mname", "rname", "serial", "refresh", "retry", "expire", "minimum"] - def __init__(self, rdclass, rdtype, mname, rname, serial, refresh, retry, - expire, minimum): + def __init__( + self, rdclass, rdtype, mname, rname, serial, refresh, retry, expire, minimum + ): super().__init__(rdclass, rdtype) self.mname = self._as_name(mname) self.rname = self._as_name(rname) @@ -47,13 +47,20 @@ class SOA(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): mname = self.mname.choose_relativity(origin, relativize) rname = self.rname.choose_relativity(origin, relativize) - return '%s %s %d %d %d %d %d' % ( - mname, rname, self.serial, self.refresh, self.retry, - self.expire, self.minimum) + return "%s %s %d %d %d %d %d" % ( + mname, + rname, + self.serial, + self.refresh, + self.retry, + self.expire, + self.minimum, + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): mname = tok.get_name(origin, relativize, relativize_to) rname = tok.get_name(origin, relativize, relativize_to) serial = tok.get_uint32() @@ -61,18 +68,20 @@ class SOA(dns.rdata.Rdata): retry = tok.get_ttl() expire = tok.get_ttl() minimum = tok.get_ttl() - return cls(rdclass, rdtype, mname, rname, serial, refresh, retry, - expire, minimum) + return cls( + rdclass, rdtype, mname, rname, serial, refresh, retry, expire, minimum + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): self.mname.to_wire(file, compress, origin, canonicalize) self.rname.to_wire(file, compress, origin, canonicalize) - five_ints = struct.pack('!IIIII', self.serial, self.refresh, - self.retry, self.expire, self.minimum) + five_ints = struct.pack( + "!IIIII", self.serial, self.refresh, self.retry, self.expire, self.minimum + ) file.write(five_ints) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): mname = parser.get_name(origin) rname = parser.get_name(origin) - return cls(rdclass, rdtype, mname, rname, *parser.get_struct('!IIIII')) + return cls(rdclass, rdtype, mname, rname, *parser.get_struct("!IIIII")) diff --git a/dns/rdtypes/ANY/SSHFP.py b/dns/rdtypes/ANY/SSHFP.py index cc03519..58ffcbb 100644 --- a/dns/rdtypes/ANY/SSHFP.py +++ b/dns/rdtypes/ANY/SSHFP.py @@ -30,10 +30,9 @@ class SSHFP(dns.rdata.Rdata): # See RFC 4255 - __slots__ = ['algorithm', 'fp_type', 'fingerprint'] + __slots__ = ["algorithm", "fp_type", "fingerprint"] - def __init__(self, rdclass, rdtype, algorithm, fp_type, - fingerprint): + def __init__(self, rdclass, rdtype, algorithm, fp_type, fingerprint): super().__init__(rdclass, rdtype) self.algorithm = self._as_uint8(algorithm) self.fp_type = self._as_uint8(fp_type) @@ -41,16 +40,17 @@ class SSHFP(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): kw = kw.copy() - chunksize = kw.pop('chunksize', 128) - return '%d %d %s' % (self.algorithm, - self.fp_type, - dns.rdata._hexify(self.fingerprint, - chunksize=chunksize, - **kw)) + chunksize = kw.pop("chunksize", 128) + return "%d %d %s" % ( + self.algorithm, + self.fp_type, + dns.rdata._hexify(self.fingerprint, chunksize=chunksize, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_uint8() fp_type = tok.get_uint8() fingerprint = tok.concatenate_remaining_identifiers().encode() diff --git a/dns/rdtypes/ANY/TKEY.py b/dns/rdtypes/ANY/TKEY.py index 59ffe03..070f03a 100644 --- a/dns/rdtypes/ANY/TKEY.py +++ b/dns/rdtypes/ANY/TKEY.py @@ -28,11 +28,28 @@ class TKEY(dns.rdata.Rdata): """TKEY Record""" - __slots__ = ['algorithm', 'inception', 'expiration', 'mode', 'error', - 'key', 'other'] - - def __init__(self, rdclass, rdtype, algorithm, inception, expiration, - mode, error, key, other=b''): + __slots__ = [ + "algorithm", + "inception", + "expiration", + "mode", + "error", + "key", + "other", + ] + + def __init__( + self, + rdclass, + rdtype, + algorithm, + inception, + expiration, + mode, + error, + key, + other=b"", + ): super().__init__(rdclass, rdtype) self.algorithm = self._as_name(algorithm) self.inception = self._as_uint32(inception) @@ -44,17 +61,23 @@ class TKEY(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): _algorithm = self.algorithm.choose_relativity(origin, relativize) - text = '%s %u %u %u %u %s' % (str(_algorithm), self.inception, - self.expiration, self.mode, self.error, - dns.rdata._base64ify(self.key, 0)) + text = "%s %u %u %u %u %s" % ( + str(_algorithm), + self.inception, + self.expiration, + self.mode, + self.error, + dns.rdata._base64ify(self.key, 0), + ) if len(self.other) > 0: - text += ' %s' % (dns.rdata._base64ify(self.other, 0)) + text += " %s" % (dns.rdata._base64ify(self.other, 0)) return text @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_name(relativize=False) inception = tok.get_uint32() expiration = tok.get_uint32() @@ -65,13 +88,15 @@ class TKEY(dns.rdata.Rdata): other_b64 = tok.concatenate_remaining_identifiers(True).encode() other = base64.b64decode(other_b64) - return cls(rdclass, rdtype, algorithm, inception, expiration, mode, - error, key, other) + return cls( + rdclass, rdtype, algorithm, inception, expiration, mode, error, key, other + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): self.algorithm.to_wire(file, compress, origin) - file.write(struct.pack("!IIHH", self.inception, self.expiration, - self.mode, self.error)) + file.write( + struct.pack("!IIHH", self.inception, self.expiration, self.mode, self.error) + ) file.write(struct.pack("!H", len(self.key))) file.write(self.key) file.write(struct.pack("!H", len(self.other))) @@ -85,8 +110,9 @@ class TKEY(dns.rdata.Rdata): key = parser.get_counted_bytes(2) other = parser.get_counted_bytes(2) - return cls(rdclass, rdtype, algorithm, inception, expiration, mode, - error, key, other) + return cls( + rdclass, rdtype, algorithm, inception, expiration, mode, error, key, other + ) # Constants for the mode field - from RFC 2930: # 2.5 The Mode Field diff --git a/dns/rdtypes/ANY/TSIG.py b/dns/rdtypes/ANY/TSIG.py index b43a78f..1ae87eb 100644 --- a/dns/rdtypes/ANY/TSIG.py +++ b/dns/rdtypes/ANY/TSIG.py @@ -29,11 +29,28 @@ class TSIG(dns.rdata.Rdata): """TSIG record""" - __slots__ = ['algorithm', 'time_signed', 'fudge', 'mac', - 'original_id', 'error', 'other'] - - def __init__(self, rdclass, rdtype, algorithm, time_signed, fudge, mac, - original_id, error, other): + __slots__ = [ + "algorithm", + "time_signed", + "fudge", + "mac", + "original_id", + "error", + "other", + ] + + def __init__( + self, + rdclass, + rdtype, + algorithm, + time_signed, + fudge, + mac, + original_id, + error, + other, + ): """Initialize a TSIG rdata. *rdclass*, an ``int`` is the rdataclass of the Rdata. @@ -67,45 +84,60 @@ class TSIG(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): algorithm = self.algorithm.choose_relativity(origin, relativize) error = dns.rcode.to_text(self.error, True) - text = f"{algorithm} {self.time_signed} {self.fudge} " + \ - f"{len(self.mac)} {dns.rdata._base64ify(self.mac, 0)} " + \ - f"{self.original_id} {error} {len(self.other)}" + text = ( + f"{algorithm} {self.time_signed} {self.fudge} " + + f"{len(self.mac)} {dns.rdata._base64ify(self.mac, 0)} " + + f"{self.original_id} {error} {len(self.other)}" + ) if self.other: text += f" {dns.rdata._base64ify(self.other, 0)}" return text @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): algorithm = tok.get_name(relativize=False) time_signed = tok.get_uint48() fudge = tok.get_uint16() mac_len = tok.get_uint16() mac = base64.b64decode(tok.get_string()) if len(mac) != mac_len: - raise SyntaxError('invalid MAC') + raise SyntaxError("invalid MAC") original_id = tok.get_uint16() error = dns.rcode.from_text(tok.get_string()) other_len = tok.get_uint16() if other_len > 0: other = base64.b64decode(tok.get_string()) if len(other) != other_len: - raise SyntaxError('invalid other data') + raise SyntaxError("invalid other data") else: - other = b'' - return cls(rdclass, rdtype, algorithm, time_signed, fudge, mac, - original_id, error, other) + other = b"" + return cls( + rdclass, + rdtype, + algorithm, + time_signed, + fudge, + mac, + original_id, + error, + other, + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): self.algorithm.to_wire(file, None, origin, False) - file.write(struct.pack('!HIHH', - (self.time_signed >> 32) & 0xffff, - self.time_signed & 0xffffffff, - self.fudge, - len(self.mac))) + file.write( + struct.pack( + "!HIHH", + (self.time_signed >> 32) & 0xFFFF, + self.time_signed & 0xFFFFFFFF, + self.fudge, + len(self.mac), + ) + ) file.write(self.mac) - file.write(struct.pack('!HHH', self.original_id, self.error, - len(self.other))) + file.write(struct.pack("!HHH", self.original_id, self.error, len(self.other))) file.write(self.other) @classmethod @@ -114,7 +146,16 @@ class TSIG(dns.rdata.Rdata): time_signed = parser.get_uint48() fudge = parser.get_uint16() mac = parser.get_counted_bytes(2) - (original_id, error) = parser.get_struct('!HH') + (original_id, error) = parser.get_struct("!HH") other = parser.get_counted_bytes(2) - return cls(rdclass, rdtype, algorithm, time_signed, fudge, mac, - original_id, error, other) + return cls( + rdclass, + rdtype, + algorithm, + time_signed, + fudge, + mac, + original_id, + error, + other, + ) diff --git a/dns/rdtypes/ANY/URI.py b/dns/rdtypes/ANY/URI.py index 524fa1b..b4c95a3 100644 --- a/dns/rdtypes/ANY/URI.py +++ b/dns/rdtypes/ANY/URI.py @@ -32,7 +32,7 @@ class URI(dns.rdata.Rdata): # see RFC 7553 - __slots__ = ['priority', 'weight', 'target'] + __slots__ = ["priority", "weight", "target"] def __init__(self, rdclass, rdtype, priority, weight, target): super().__init__(rdclass, rdtype) @@ -43,12 +43,12 @@ class URI(dns.rdata.Rdata): raise dns.exception.SyntaxError("URI target cannot be empty") def to_text(self, origin=None, relativize=True, **kw): - return '%d %d "%s"' % (self.priority, self.weight, - self.target.decode()) + return '%d %d "%s"' % (self.priority, self.weight, self.target.decode()) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): priority = tok.get_uint16() weight = tok.get_uint16() target = tok.get().unescape() @@ -63,10 +63,10 @@ class URI(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (priority, weight) = parser.get_struct('!HH') + (priority, weight) = parser.get_struct("!HH") target = parser.get_remaining() if len(target) == 0: - raise dns.exception.FormError('URI target may not be empty') + raise dns.exception.FormError("URI target may not be empty") return cls(rdclass, rdtype, priority, weight, target) def _processing_priority(self): diff --git a/dns/rdtypes/ANY/X25.py b/dns/rdtypes/ANY/X25.py index 4f7230c..06c1453 100644 --- a/dns/rdtypes/ANY/X25.py +++ b/dns/rdtypes/ANY/X25.py @@ -30,7 +30,7 @@ class X25(dns.rdata.Rdata): # see RFC 1183 - __slots__ = ['address'] + __slots__ = ["address"] def __init__(self, rdclass, rdtype, address): super().__init__(rdclass, rdtype) @@ -40,15 +40,16 @@ class X25(dns.rdata.Rdata): return '"%s"' % dns.rdata._escapify(self.address) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_string() return cls(rdclass, rdtype, address) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): l = len(self.address) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(self.address) @classmethod diff --git a/dns/rdtypes/ANY/ZONEMD.py b/dns/rdtypes/ANY/ZONEMD.py index 75f99e5..1f86ba4 100644 --- a/dns/rdtypes/ANY/ZONEMD.py +++ b/dns/rdtypes/ANY/ZONEMD.py @@ -16,7 +16,7 @@ class ZONEMD(dns.rdata.Rdata): # See RFC 8976 - __slots__ = ['serial', 'scheme', 'hash_algorithm', 'digest'] + __slots__ = ["serial", "scheme", "hash_algorithm", "digest"] def __init__(self, rdclass, rdtype, serial, scheme, hash_algorithm, digest): super().__init__(rdclass, rdtype) @@ -26,25 +26,28 @@ class ZONEMD(dns.rdata.Rdata): self.digest = self._as_bytes(digest) if self.scheme == 0: # reserved, RFC 8976 Sec. 5.2 - raise ValueError('scheme 0 is reserved') + raise ValueError("scheme 0 is reserved") if self.hash_algorithm == 0: # reserved, RFC 8976 Sec. 5.3 - raise ValueError('hash_algorithm 0 is reserved') + raise ValueError("hash_algorithm 0 is reserved") hasher = dns.zonetypes._digest_hashers.get(self.hash_algorithm) if hasher and hasher().digest_size != len(self.digest): - raise ValueError('digest length inconsistent with hash algorithm') + raise ValueError("digest length inconsistent with hash algorithm") def to_text(self, origin=None, relativize=True, **kw): kw = kw.copy() - chunksize = kw.pop('chunksize', 128) - return '%d %d %d %s' % (self.serial, self.scheme, self.hash_algorithm, - dns.rdata._hexify(self.digest, - chunksize=chunksize, - **kw)) + chunksize = kw.pop("chunksize", 128) + return "%d %d %d %s" % ( + self.serial, + self.scheme, + self.hash_algorithm, + dns.rdata._hexify(self.digest, chunksize=chunksize, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): serial = tok.get_uint32() scheme = tok.get_uint8() hash_algorithm = tok.get_uint8() @@ -53,8 +56,7 @@ class ZONEMD(dns.rdata.Rdata): return cls(rdclass, rdtype, serial, scheme, hash_algorithm, digest) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - header = struct.pack("!IBB", self.serial, self.scheme, - self.hash_algorithm) + header = struct.pack("!IBB", self.serial, self.scheme, self.hash_algorithm) file.write(header) file.write(self.digest) diff --git a/dns/rdtypes/ANY/__init__.py b/dns/rdtypes/ANY/__init__.py index 2cadcde..3824a0a 100644 --- a/dns/rdtypes/ANY/__init__.py +++ b/dns/rdtypes/ANY/__init__.py @@ -18,51 +18,51 @@ """Class ANY (generic) rdata type classes.""" __all__ = [ - 'AFSDB', - 'AMTRELAY', - 'AVC', - 'CAA', - 'CDNSKEY', - 'CDS', - 'CERT', - 'CNAME', - 'CSYNC', - 'DLV', - 'DNAME', - 'DNSKEY', - 'DS', - 'EUI48', - 'EUI64', - 'GPOS', - 'HINFO', - 'HIP', - 'ISDN', - 'L32', - 'L64', - 'LOC', - 'LP', - 'MX', - 'NID', - 'NINFO', - 'NS', - 'NSEC', - 'NSEC3', - 'NSEC3PARAM', - 'OPENPGPKEY', - 'OPT', - 'PTR', - 'RP', - 'RRSIG', - 'RT', - 'SMIMEA', - 'SOA', - 'SPF', - 'SSHFP', - 'TKEY', - 'TLSA', - 'TSIG', - 'TXT', - 'URI', - 'X25', - 'ZONEMD', + "AFSDB", + "AMTRELAY", + "AVC", + "CAA", + "CDNSKEY", + "CDS", + "CERT", + "CNAME", + "CSYNC", + "DLV", + "DNAME", + "DNSKEY", + "DS", + "EUI48", + "EUI64", + "GPOS", + "HINFO", + "HIP", + "ISDN", + "L32", + "L64", + "LOC", + "LP", + "MX", + "NID", + "NINFO", + "NS", + "NSEC", + "NSEC3", + "NSEC3PARAM", + "OPENPGPKEY", + "OPT", + "PTR", + "RP", + "RRSIG", + "RT", + "SMIMEA", + "SOA", + "SPF", + "SSHFP", + "TKEY", + "TLSA", + "TSIG", + "TXT", + "URI", + "X25", + "ZONEMD", ] diff --git a/dns/rdtypes/CH/A.py b/dns/rdtypes/CH/A.py index 828701b..9905c7c 100644 --- a/dns/rdtypes/CH/A.py +++ b/dns/rdtypes/CH/A.py @@ -20,6 +20,7 @@ import struct import dns.rdtypes.mxbase import dns.immutable + @dns.immutable.immutable class A(dns.rdata.Rdata): @@ -28,7 +29,7 @@ class A(dns.rdata.Rdata): # domain: the domain of the address # address: the 16-bit address - __slots__ = ['domain', 'address'] + __slots__ = ["domain", "address"] def __init__(self, rdclass, rdtype, domain, address): super().__init__(rdclass, rdtype) @@ -37,11 +38,12 @@ class A(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): domain = self.domain.choose_relativity(origin, relativize) - return '%s %o' % (domain, self.address) + return "%s %o" % (domain, self.address) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): domain = tok.get_name(origin, relativize, relativize_to) address = tok.get_uint16(base=8) return cls(rdclass, rdtype, domain, address) diff --git a/dns/rdtypes/CH/__init__.py b/dns/rdtypes/CH/__init__.py index 7184a73..0760c26 100644 --- a/dns/rdtypes/CH/__init__.py +++ b/dns/rdtypes/CH/__init__.py @@ -18,5 +18,5 @@ """Class CH rdata type classes.""" __all__ = [ - 'A', + "A", ] diff --git a/dns/rdtypes/IN/A.py b/dns/rdtypes/IN/A.py index 74b591e..713d5ee 100644 --- a/dns/rdtypes/IN/A.py +++ b/dns/rdtypes/IN/A.py @@ -27,7 +27,7 @@ class A(dns.rdata.Rdata): """A record.""" - __slots__ = ['address'] + __slots__ = ["address"] def __init__(self, rdclass, rdtype, address): super().__init__(rdclass, rdtype) @@ -37,8 +37,9 @@ class A(dns.rdata.Rdata): return self.address @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_identifier() return cls(rdclass, rdtype, address) diff --git a/dns/rdtypes/IN/AAAA.py b/dns/rdtypes/IN/AAAA.py index 2d3ec90..f8237b4 100644 --- a/dns/rdtypes/IN/AAAA.py +++ b/dns/rdtypes/IN/AAAA.py @@ -27,7 +27,7 @@ class AAAA(dns.rdata.Rdata): """AAAA record.""" - __slots__ = ['address'] + __slots__ = ["address"] def __init__(self, rdclass, rdtype, address): super().__init__(rdclass, rdtype) @@ -37,8 +37,9 @@ class AAAA(dns.rdata.Rdata): return self.address @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_identifier() return cls(rdclass, rdtype, address) diff --git a/dns/rdtypes/IN/APL.py b/dns/rdtypes/IN/APL.py index ae94fb2..05e1689 100644 --- a/dns/rdtypes/IN/APL.py +++ b/dns/rdtypes/IN/APL.py @@ -26,12 +26,13 @@ import dns.ipv6 import dns.rdata import dns.tokenizer + @dns.immutable.immutable class APLItem: """An APL list item.""" - __slots__ = ['family', 'negation', 'address', 'prefix'] + __slots__ = ["family", "negation", "address", "prefix"] def __init__(self, family, negation, address, prefix): self.family = dns.rdata.Rdata._as_uint16(family) @@ -67,12 +68,12 @@ class APLItem: if address[i] != 0: last = i + 1 break - address = address[0: last] + address = address[0:last] l = len(address) assert l < 128 if self.negation: l |= 0x80 - header = struct.pack('!HBB', self.family, self.prefix, l) + header = struct.pack("!HBB", self.family, self.prefix, l) file.write(header) file.write(address) @@ -84,32 +85,33 @@ class APL(dns.rdata.Rdata): # see: RFC 3123 - __slots__ = ['items'] + __slots__ = ["items"] def __init__(self, rdclass, rdtype, items): super().__init__(rdclass, rdtype) for item in items: if not isinstance(item, APLItem): - raise ValueError('item not an APLItem') + raise ValueError("item not an APLItem") self.items = tuple(items) def to_text(self, origin=None, relativize=True, **kw): - return ' '.join(map(str, self.items)) + return " ".join(map(str, self.items)) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): items = [] for token in tok.get_remaining(): item = token.unescape().value - if item[0] == '!': + if item[0] == "!": negation = True item = item[1:] else: negation = False - (family, rest) = item.split(':', 1) + (family, rest) = item.split(":", 1) family = int(family) - (address, prefix) = rest.split('/', 1) + (address, prefix) = rest.split("/", 1) prefix = int(prefix) item = APLItem(family, negation, address, prefix) items.append(item) @@ -125,7 +127,7 @@ class APL(dns.rdata.Rdata): items = [] while parser.remaining() > 0: - header = parser.get_struct('!HBB') + header = parser.get_struct("!HBB") afdlen = header[2] if afdlen > 127: negation = True @@ -136,16 +138,16 @@ class APL(dns.rdata.Rdata): l = len(address) if header[0] == 1: if l < 4: - address += b'\x00' * (4 - l) + address += b"\x00" * (4 - l) elif header[0] == 2: if l < 16: - address += b'\x00' * (16 - l) + address += b"\x00" * (16 - l) else: # # This isn't really right according to the RFC, but it # seems better than throwing an exception # - address = codecs.encode(address, 'hex_codec') + address = codecs.encode(address, "hex_codec") item = APLItem(header[0], negation, address, header[1]) items.append(item) return cls(rdclass, rdtype, items) diff --git a/dns/rdtypes/IN/DHCID.py b/dns/rdtypes/IN/DHCID.py index c1c70b4..65f8589 100644 --- a/dns/rdtypes/IN/DHCID.py +++ b/dns/rdtypes/IN/DHCID.py @@ -29,7 +29,7 @@ class DHCID(dns.rdata.Rdata): # see: RFC 4701 - __slots__ = ['data'] + __slots__ = ["data"] def __init__(self, rdclass, rdtype, data): super().__init__(rdclass, rdtype) @@ -39,8 +39,9 @@ class DHCID(dns.rdata.Rdata): return dns.rdata._base64ify(self.data, **kw) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): b64 = tok.concatenate_remaining_identifiers().encode() data = base64.b64decode(b64) return cls(rdclass, rdtype, data) diff --git a/dns/rdtypes/IN/HTTPS.py b/dns/rdtypes/IN/HTTPS.py index 6a67e8e..7797fba 100644 --- a/dns/rdtypes/IN/HTTPS.py +++ b/dns/rdtypes/IN/HTTPS.py @@ -3,6 +3,7 @@ import dns.rdtypes.svcbbase import dns.immutable + @dns.immutable.immutable class HTTPS(dns.rdtypes.svcbbase.SVCBBase): """HTTPS record""" diff --git a/dns/rdtypes/IN/IPSECKEY.py b/dns/rdtypes/IN/IPSECKEY.py index d1d3943..1255739 100644 --- a/dns/rdtypes/IN/IPSECKEY.py +++ b/dns/rdtypes/IN/IPSECKEY.py @@ -24,7 +24,8 @@ import dns.rdtypes.util class Gateway(dns.rdtypes.util.Gateway): - name = 'IPSECKEY gateway' + name = "IPSECKEY gateway" + @dns.immutable.immutable class IPSECKEY(dns.rdata.Rdata): @@ -33,10 +34,11 @@ class IPSECKEY(dns.rdata.Rdata): # see: RFC 4025 - __slots__ = ['precedence', 'gateway_type', 'algorithm', 'gateway', 'key'] + __slots__ = ["precedence", "gateway_type", "algorithm", "gateway", "key"] - def __init__(self, rdclass, rdtype, precedence, gateway_type, algorithm, - gateway, key): + def __init__( + self, rdclass, rdtype, precedence, gateway_type, algorithm, gateway, key + ): super().__init__(rdclass, rdtype) gateway = Gateway(gateway_type, gateway) self.precedence = self._as_uint8(precedence) @@ -46,38 +48,45 @@ class IPSECKEY(dns.rdata.Rdata): self.key = self._as_bytes(key) def to_text(self, origin=None, relativize=True, **kw): - gateway = Gateway(self.gateway_type, self.gateway).to_text(origin, - relativize) - return '%d %d %d %s %s' % (self.precedence, self.gateway_type, - self.algorithm, gateway, - dns.rdata._base64ify(self.key, **kw)) + gateway = Gateway(self.gateway_type, self.gateway).to_text(origin, relativize) + return "%d %d %d %s %s" % ( + self.precedence, + self.gateway_type, + self.algorithm, + gateway, + dns.rdata._base64ify(self.key, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): precedence = tok.get_uint8() gateway_type = tok.get_uint8() algorithm = tok.get_uint8() - gateway = Gateway.from_text(gateway_type, tok, origin, relativize, - relativize_to) + gateway = Gateway.from_text( + gateway_type, tok, origin, relativize, relativize_to + ) b64 = tok.concatenate_remaining_identifiers().encode() key = base64.b64decode(b64) - return cls(rdclass, rdtype, precedence, gateway_type, algorithm, - gateway.gateway, key) + return cls( + rdclass, rdtype, precedence, gateway_type, algorithm, gateway.gateway, key + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - header = struct.pack("!BBB", self.precedence, self.gateway_type, - self.algorithm) + header = struct.pack("!BBB", self.precedence, self.gateway_type, self.algorithm) file.write(header) - Gateway(self.gateway_type, self.gateway).to_wire(file, compress, - origin, canonicalize) + Gateway(self.gateway_type, self.gateway).to_wire( + file, compress, origin, canonicalize + ) file.write(self.key) @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - header = parser.get_struct('!BBB') + header = parser.get_struct("!BBB") gateway_type = header[1] gateway = Gateway.from_wire_parser(gateway_type, parser, origin) key = parser.get_remaining() - return cls(rdclass, rdtype, header[0], gateway_type, header[2], - gateway.gateway, key) + return cls( + rdclass, rdtype, header[0], gateway_type, header[2], gateway.gateway, key + ) diff --git a/dns/rdtypes/IN/NAPTR.py b/dns/rdtypes/IN/NAPTR.py index b107974..1f1f5a1 100644 --- a/dns/rdtypes/IN/NAPTR.py +++ b/dns/rdtypes/IN/NAPTR.py @@ -27,7 +27,7 @@ import dns.rdtypes.util def _write_string(file, s): l = len(s) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(s) @@ -38,11 +38,11 @@ class NAPTR(dns.rdata.Rdata): # see: RFC 3403 - __slots__ = ['order', 'preference', 'flags', 'service', 'regexp', - 'replacement'] + __slots__ = ["order", "preference", "flags", "service", "regexp", "replacement"] - def __init__(self, rdclass, rdtype, order, preference, flags, service, - regexp, replacement): + def __init__( + self, rdclass, rdtype, order, preference, flags, service, regexp, replacement + ): super().__init__(rdclass, rdtype) self.flags = self._as_bytes(flags, True, 255) self.service = self._as_bytes(service, True, 255) @@ -53,24 +53,28 @@ class NAPTR(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): replacement = self.replacement.choose_relativity(origin, relativize) - return '%d %d "%s" "%s" "%s" %s' % \ - (self.order, self.preference, - dns.rdata._escapify(self.flags), - dns.rdata._escapify(self.service), - dns.rdata._escapify(self.regexp), - replacement) + return '%d %d "%s" "%s" "%s" %s' % ( + self.order, + self.preference, + dns.rdata._escapify(self.flags), + dns.rdata._escapify(self.service), + dns.rdata._escapify(self.regexp), + replacement, + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): order = tok.get_uint16() preference = tok.get_uint16() flags = tok.get_string() service = tok.get_string() regexp = tok.get_string() replacement = tok.get_name(origin, relativize, relativize_to) - return cls(rdclass, rdtype, order, preference, flags, service, - regexp, replacement) + return cls( + rdclass, rdtype, order, preference, flags, service, regexp, replacement + ) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): two_ints = struct.pack("!HH", self.order, self.preference) @@ -82,14 +86,22 @@ class NAPTR(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (order, preference) = parser.get_struct('!HH') + (order, preference) = parser.get_struct("!HH") strings = [] for _ in range(3): s = parser.get_counted_bytes() strings.append(s) replacement = parser.get_name(origin) - return cls(rdclass, rdtype, order, preference, strings[0], strings[1], - strings[2], replacement) + return cls( + rdclass, + rdtype, + order, + preference, + strings[0], + strings[1], + strings[2], + replacement, + ) def _processing_priority(self): return (self.order, self.preference) diff --git a/dns/rdtypes/IN/NSAP.py b/dns/rdtypes/IN/NSAP.py index 23ae9b1..be8581e 100644 --- a/dns/rdtypes/IN/NSAP.py +++ b/dns/rdtypes/IN/NSAP.py @@ -30,7 +30,7 @@ class NSAP(dns.rdata.Rdata): # see: RFC 1706 - __slots__ = ['address'] + __slots__ = ["address"] def __init__(self, rdclass, rdtype, address): super().__init__(rdclass, rdtype) @@ -40,14 +40,15 @@ class NSAP(dns.rdata.Rdata): return "0x%s" % binascii.hexlify(self.address).decode() @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_string() - if address[0:2] != '0x': - raise dns.exception.SyntaxError('string does not start with 0x') - address = address[2:].replace('.', '') + if address[0:2] != "0x": + raise dns.exception.SyntaxError("string does not start with 0x") + address = address[2:].replace(".", "") if len(address) % 2 != 0: - raise dns.exception.SyntaxError('hexstring has odd length') + raise dns.exception.SyntaxError("hexstring has odd length") address = binascii.unhexlify(address.encode()) return cls(rdclass, rdtype, address) diff --git a/dns/rdtypes/IN/PX.py b/dns/rdtypes/IN/PX.py index 113d409..b2216d6 100644 --- a/dns/rdtypes/IN/PX.py +++ b/dns/rdtypes/IN/PX.py @@ -31,7 +31,7 @@ class PX(dns.rdata.Rdata): # see: RFC 2163 - __slots__ = ['preference', 'map822', 'mapx400'] + __slots__ = ["preference", "map822", "mapx400"] def __init__(self, rdclass, rdtype, preference, map822, mapx400): super().__init__(rdclass, rdtype) @@ -42,11 +42,12 @@ class PX(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): map822 = self.map822.choose_relativity(origin, relativize) mapx400 = self.mapx400.choose_relativity(origin, relativize) - return '%d %s %s' % (self.preference, map822, mapx400) + return "%d %s %s" % (self.preference, map822, mapx400) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() map822 = tok.get_name(origin, relativize, relativize_to) mapx400 = tok.get_name(origin, relativize, relativize_to) diff --git a/dns/rdtypes/IN/SRV.py b/dns/rdtypes/IN/SRV.py index 5b5ff42..8b0b6bf 100644 --- a/dns/rdtypes/IN/SRV.py +++ b/dns/rdtypes/IN/SRV.py @@ -31,7 +31,7 @@ class SRV(dns.rdata.Rdata): # see: RFC 2782 - __slots__ = ['priority', 'weight', 'port', 'target'] + __slots__ = ["priority", "weight", "port", "target"] def __init__(self, rdclass, rdtype, priority, weight, port, target): super().__init__(rdclass, rdtype) @@ -42,12 +42,12 @@ class SRV(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): target = self.target.choose_relativity(origin, relativize) - return '%d %d %d %s' % (self.priority, self.weight, self.port, - target) + return "%d %d %d %s" % (self.priority, self.weight, self.port, target) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): priority = tok.get_uint16() weight = tok.get_uint16() port = tok.get_uint16() @@ -61,7 +61,7 @@ class SRV(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - (priority, weight, port) = parser.get_struct('!HHH') + (priority, weight, port) = parser.get_struct("!HHH") target = parser.get_name(origin) return cls(rdclass, rdtype, priority, weight, port, target) diff --git a/dns/rdtypes/IN/SVCB.py b/dns/rdtypes/IN/SVCB.py index 14838e1..9a1ad10 100644 --- a/dns/rdtypes/IN/SVCB.py +++ b/dns/rdtypes/IN/SVCB.py @@ -3,6 +3,7 @@ import dns.rdtypes.svcbbase import dns.immutable + @dns.immutable.immutable class SVCB(dns.rdtypes.svcbbase.SVCBBase): """SVCB record""" diff --git a/dns/rdtypes/IN/WKS.py b/dns/rdtypes/IN/WKS.py index 264e45d..a671e20 100644 --- a/dns/rdtypes/IN/WKS.py +++ b/dns/rdtypes/IN/WKS.py @@ -23,13 +23,14 @@ import dns.immutable import dns.rdata try: - _proto_tcp = socket.getprotobyname('tcp') - _proto_udp = socket.getprotobyname('udp') + _proto_tcp = socket.getprotobyname("tcp") + _proto_udp = socket.getprotobyname("udp") except OSError: # Fall back to defaults in case /etc/protocols is unavailable. _proto_tcp = 6 _proto_udp = 17 + @dns.immutable.immutable class WKS(dns.rdata.Rdata): @@ -37,7 +38,7 @@ class WKS(dns.rdata.Rdata): # see: RFC 1035 - __slots__ = ['address', 'protocol', 'bitmap'] + __slots__ = ["address", "protocol", "bitmap"] def __init__(self, rdclass, rdtype, address, protocol, bitmap): super().__init__(rdclass, rdtype) @@ -51,12 +52,13 @@ class WKS(dns.rdata.Rdata): for j in range(0, 8): if byte & (0x80 >> j): bits.append(str(i * 8 + j)) - text = ' '.join(bits) - return '%s %d %s' % (self.address, self.protocol, text) + text = " ".join(bits) + return "%s %d %s" % (self.address, self.protocol, text) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): address = tok.get_string() protocol = tok.get_string() if protocol.isdigit(): @@ -87,7 +89,7 @@ class WKS(dns.rdata.Rdata): def _to_wire(self, file, compress=None, origin=None, canonicalize=False): file.write(dns.ipv4.inet_aton(self.address)) - protocol = struct.pack('!B', self.protocol) + protocol = struct.pack("!B", self.protocol) file.write(protocol) file.write(self.bitmap) diff --git a/dns/rdtypes/IN/__init__.py b/dns/rdtypes/IN/__init__.py index d51b99e..dcec4dd 100644 --- a/dns/rdtypes/IN/__init__.py +++ b/dns/rdtypes/IN/__init__.py @@ -18,18 +18,18 @@ """Class IN rdata type classes.""" __all__ = [ - 'A', - 'AAAA', - 'APL', - 'DHCID', - 'HTTPS', - 'IPSECKEY', - 'KX', - 'NAPTR', - 'NSAP', - 'NSAP_PTR', - 'PX', - 'SRV', - 'SVCB', - 'WKS', + "A", + "AAAA", + "APL", + "DHCID", + "HTTPS", + "IPSECKEY", + "KX", + "NAPTR", + "NSAP", + "NSAP_PTR", + "PX", + "SRV", + "SVCB", + "WKS", ] diff --git a/dns/rdtypes/__init__.py b/dns/rdtypes/__init__.py index c3af264..3997f84 100644 --- a/dns/rdtypes/__init__.py +++ b/dns/rdtypes/__init__.py @@ -18,16 +18,16 @@ """DNS rdata type classes""" __all__ = [ - 'ANY', - 'IN', - 'CH', - 'dnskeybase', - 'dsbase', - 'euibase', - 'mxbase', - 'nsbase', - 'svcbbase', - 'tlsabase', - 'txtbase', - 'util' + "ANY", + "IN", + "CH", + "dnskeybase", + "dsbase", + "euibase", + "mxbase", + "nsbase", + "svcbbase", + "tlsabase", + "txtbase", + "util", ] diff --git a/dns/rdtypes/dnskeybase.py b/dns/rdtypes/dnskeybase.py index 832df2d..1d17f70 100644 --- a/dns/rdtypes/dnskeybase.py +++ b/dns/rdtypes/dnskeybase.py @@ -25,7 +25,8 @@ import dns.dnssectypes import dns.rdata # wildcard import -__all__ = ["SEP", "REVOKE", "ZONE"] # noqa: F822 +__all__ = ["SEP", "REVOKE", "ZONE"] # noqa: F822 + class Flag(enum.IntFlag): SEP = 0x0001 @@ -38,7 +39,7 @@ class DNSKEYBase(dns.rdata.Rdata): """Base class for rdata that is like a DNSKEY record""" - __slots__ = ['flags', 'protocol', 'algorithm', 'key'] + __slots__ = ["flags", "protocol", "algorithm", "key"] def __init__(self, rdclass, rdtype, flags, protocol, algorithm, key): super().__init__(rdclass, rdtype) @@ -48,12 +49,17 @@ class DNSKEYBase(dns.rdata.Rdata): self.key = self._as_bytes(key) def to_text(self, origin=None, relativize=True, **kw): - return '%d %d %d %s' % (self.flags, self.protocol, self.algorithm, - dns.rdata._base64ify(self.key, **kw)) + return "%d %d %d %s" % ( + self.flags, + self.protocol, + self.algorithm, + dns.rdata._base64ify(self.key, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): flags = tok.get_uint16() protocol = tok.get_uint8() algorithm = tok.get_string() @@ -68,10 +74,10 @@ class DNSKEYBase(dns.rdata.Rdata): @classmethod def from_wire_parser(cls, rdclass, rdtype, parser, origin=None): - header = parser.get_struct('!HBB') + header = parser.get_struct("!HBB") key = parser.get_remaining() - return cls(rdclass, rdtype, header[0], header[1], header[2], - key) + return cls(rdclass, rdtype, header[0], header[1], header[2], key) + ### BEGIN generated Flag constants diff --git a/dns/rdtypes/dsbase.py b/dns/rdtypes/dsbase.py index 3bf93ac..b6032b0 100644 --- a/dns/rdtypes/dsbase.py +++ b/dns/rdtypes/dsbase.py @@ -29,9 +29,10 @@ class DSBase(dns.rdata.Rdata): """Base class for rdata that is like a DS record""" - __slots__ = ['key_tag', 'algorithm', 'digest_type', 'digest'] + __slots__ = ["key_tag", "algorithm", "digest_type", "digest"] - # Digest types registry: https://www.iana.org/assignments/ds-rr-types/ds-rr-types.xhtml + # Digest types registry: + # https://www.iana.org/assignments/ds-rr-types/ds-rr-types.xhtml _digest_length_by_type = { 1: 20, # SHA-1, RFC 3658 Sec. 2.4 2: 32, # SHA-256, RFC 4509 Sec. 2.2 @@ -39,8 +40,7 @@ class DSBase(dns.rdata.Rdata): 4: 48, # SHA-384, RFC 6605 Sec. 2 } - def __init__(self, rdclass, rdtype, key_tag, algorithm, digest_type, - digest): + def __init__(self, rdclass, rdtype, key_tag, algorithm, digest_type, digest): super().__init__(rdclass, rdtype) self.key_tag = self._as_uint16(key_tag) self.algorithm = dns.dnssectypes.Algorithm.make(algorithm) @@ -48,34 +48,34 @@ class DSBase(dns.rdata.Rdata): self.digest = self._as_bytes(digest) try: if len(self.digest) != self._digest_length_by_type[self.digest_type]: - raise ValueError('digest length inconsistent with digest type') + raise ValueError("digest length inconsistent with digest type") except KeyError: if self.digest_type == 0: # reserved, RFC 3658 Sec. 2.4 - raise ValueError('digest type 0 is reserved') + raise ValueError("digest type 0 is reserved") def to_text(self, origin=None, relativize=True, **kw): kw = kw.copy() - chunksize = kw.pop('chunksize', 128) - return '%d %d %d %s' % (self.key_tag, self.algorithm, - self.digest_type, - dns.rdata._hexify(self.digest, - chunksize=chunksize, - **kw)) + chunksize = kw.pop("chunksize", 128) + return "%d %d %d %s" % ( + self.key_tag, + self.algorithm, + self.digest_type, + dns.rdata._hexify(self.digest, chunksize=chunksize, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): key_tag = tok.get_uint16() algorithm = tok.get_string() digest_type = tok.get_uint8() digest = tok.concatenate_remaining_identifiers().encode() digest = binascii.unhexlify(digest) - return cls(rdclass, rdtype, key_tag, algorithm, digest_type, - digest) + return cls(rdclass, rdtype, key_tag, algorithm, digest_type, digest) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): - header = struct.pack("!HBB", self.key_tag, self.algorithm, - self.digest_type) + header = struct.pack("!HBB", self.key_tag, self.algorithm, self.digest_type) file.write(header) file.write(self.digest) diff --git a/dns/rdtypes/euibase.py b/dns/rdtypes/euibase.py index 48b69bd..e524aea 100644 --- a/dns/rdtypes/euibase.py +++ b/dns/rdtypes/euibase.py @@ -27,7 +27,7 @@ class EUIBase(dns.rdata.Rdata): # see: rfc7043.txt - __slots__ = ['eui'] + __slots__ = ["eui"] # define these in subclasses # byte_len = 6 # 0123456789ab (in hex) # text_len = byte_len * 3 - 1 # 01-23-45-67-89-ab @@ -36,28 +36,30 @@ class EUIBase(dns.rdata.Rdata): super().__init__(rdclass, rdtype) self.eui = self._as_bytes(eui) if len(self.eui) != self.byte_len: - raise dns.exception.FormError('EUI%s rdata has to have %s bytes' - % (self.byte_len * 8, self.byte_len)) + raise dns.exception.FormError( + "EUI%s rdata has to have %s bytes" % (self.byte_len * 8, self.byte_len) + ) def to_text(self, origin=None, relativize=True, **kw): - return dns.rdata._hexify(self.eui, chunksize=2, separator=b'-', **kw) + return dns.rdata._hexify(self.eui, chunksize=2, separator=b"-", **kw) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): text = tok.get_string() if len(text) != cls.text_len: raise dns.exception.SyntaxError( - 'Input text must have %s characters' % cls.text_len) + "Input text must have %s characters" % cls.text_len + ) for i in range(2, cls.byte_len * 3 - 1, 3): - if text[i] != '-': - raise dns.exception.SyntaxError('Dash expected at position %s' - % i) - text = text.replace('-', '') + if text[i] != "-": + raise dns.exception.SyntaxError("Dash expected at position %s" % i) + text = text.replace("-", "") try: data = binascii.unhexlify(text.encode()) except (ValueError, TypeError) as ex: - raise dns.exception.SyntaxError('Hex decoding error: %s' % str(ex)) + raise dns.exception.SyntaxError("Hex decoding error: %s" % str(ex)) return cls(rdclass, rdtype, data) def _to_wire(self, file, compress=None, origin=None, canonicalize=False): diff --git a/dns/rdtypes/mxbase.py b/dns/rdtypes/mxbase.py index 5641823..b4b9b08 100644 --- a/dns/rdtypes/mxbase.py +++ b/dns/rdtypes/mxbase.py @@ -31,7 +31,7 @@ class MXBase(dns.rdata.Rdata): """Base class for rdata that is like an MX record.""" - __slots__ = ['preference', 'exchange'] + __slots__ = ["preference", "exchange"] def __init__(self, rdclass, rdtype, preference, exchange): super().__init__(rdclass, rdtype) @@ -40,11 +40,12 @@ class MXBase(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): exchange = self.exchange.choose_relativity(origin, relativize) - return '%d %s' % (self.preference, exchange) + return "%d %s" % (self.preference, exchange) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): preference = tok.get_uint16() exchange = tok.get_name(origin, relativize, relativize_to) return cls(rdclass, rdtype, preference, exchange) diff --git a/dns/rdtypes/nsbase.py b/dns/rdtypes/nsbase.py index b3e2550..ba7a2ab 100644 --- a/dns/rdtypes/nsbase.py +++ b/dns/rdtypes/nsbase.py @@ -28,7 +28,7 @@ class NSBase(dns.rdata.Rdata): """Base class for rdata that is like an NS record.""" - __slots__ = ['target'] + __slots__ = ["target"] def __init__(self, rdclass, rdtype, target): super().__init__(rdclass, rdtype) @@ -39,8 +39,9 @@ class NSBase(dns.rdata.Rdata): return str(target) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): target = tok.get_name(origin, relativize, relativize_to) return cls(rdclass, rdtype, target) diff --git a/dns/rdtypes/svcbbase.py b/dns/rdtypes/svcbbase.py index d287499..a7bc273 100644 --- a/dns/rdtypes/svcbbase.py +++ b/dns/rdtypes/svcbbase.py @@ -63,44 +63,48 @@ def _validate_key(key): if isinstance(key, bytes): # We decode to latin-1 so we get 0-255 as valid and do NOT interpret # UTF-8 sequences - key = key.decode('latin-1') + key = key.decode("latin-1") if isinstance(key, str): - if key.lower().startswith('key'): + if key.lower().startswith("key"): force_generic = True - if key[3:].startswith('0') and len(key) != 4: + if key[3:].startswith("0") and len(key) != 4: # key has leading zeros - raise ValueError('leading zeros in key') - key = key.replace('-', '_') + raise ValueError("leading zeros in key") + key = key.replace("-", "_") return (ParamKey.make(key), force_generic) + def key_to_text(key): - return ParamKey.to_text(key).replace('_', '-').lower() + return ParamKey.to_text(key).replace("_", "-").lower() + # Like rdata escapify, but escapes ',' too. _escaped = b'",\\' + def _escapify(qstring): - text = '' + text = "" for c in qstring: if c in _escaped: - text += '\\' + chr(c) + text += "\\" + chr(c) elif c >= 0x20 and c < 0x7F: text += chr(c) else: - text += '\\%03d' % c + text += "\\%03d" % c return text + def _unescape(value): - if value == '': + if value == "": return value - unescaped = b'' + unescaped = b"" l = len(value) i = 0 while i < l: c = value[i] i += 1 - if c == '\\': + if c == "\\": if i >= l: # pragma: no cover (can't happen via tokenizer get()) raise dns.exception.UnexpectedEnd c = value[i] @@ -119,7 +123,7 @@ def _unescape(value): codepoint = int(c) * 100 + int(c2) * 10 + int(c3) if codepoint > 255: raise dns.exception.SyntaxError - unescaped += b'%c' % (codepoint) + unescaped += b"%c" % (codepoint) continue unescaped += c.encode() return unescaped @@ -129,21 +133,21 @@ def _split(value): l = len(value) i = 0 items = [] - unescaped = b'' + unescaped = b"" while i < l: c = value[i] i += 1 - if c == ord('\\'): + if c == ord("\\"): if i >= l: # pragma: no cover (can't happen via tokenizer get()) raise dns.exception.UnexpectedEnd c = value[i] i += 1 - unescaped += b'%c' % (c) - elif c == ord(','): + unescaped += b"%c" % (c) + elif c == ord(","): items.append(unescaped) - unescaped = b'' + unescaped = b"" else: - unescaped += b'%c' % (c) + unescaped += b"%c" % (c) items.append(unescaped) return items @@ -159,8 +163,8 @@ class Param: @dns.immutable.immutable class GenericParam(Param): - """Generic SVCB parameter - """ + """Generic SVCB parameter""" + def __init__(self, value): self.value = dns.rdata.Rdata._as_bytes(value, True) @@ -198,19 +202,19 @@ class MandatoryParam(Param): prior_k = None for k in keys: if k == prior_k: - raise ValueError(f'duplicate key {k:d}') + raise ValueError(f"duplicate key {k:d}") prior_k = k if k == ParamKey.MANDATORY: - raise ValueError('listed the mandatory key as mandatory') + raise ValueError("listed the mandatory key as mandatory") self.keys = tuple(keys) @classmethod def from_value(cls, value): - keys = [k.encode() for k in value.split(',')] + keys = [k.encode() for k in value.split(",")] return cls(keys) def to_text(self): - return '"' + ','.join([key_to_text(key) for key in self.keys]) + '"' + return '"' + ",".join([key_to_text(key) for key in self.keys]) + '"' @classmethod def from_wire_parser(cls, parser, origin=None): # pylint: disable=W0613 @@ -219,28 +223,29 @@ class MandatoryParam(Param): while parser.remaining() > 0: key = parser.get_uint16() if key < last_key: - raise dns.exception.FormError('manadatory keys not ascending') + raise dns.exception.FormError("manadatory keys not ascending") last_key = key keys.append(key) return cls(keys) def to_wire(self, file, origin=None): # pylint: disable=W0613 for key in self.keys: - file.write(struct.pack('!H', key)) + file.write(struct.pack("!H", key)) @dns.immutable.immutable class ALPNParam(Param): def __init__(self, ids): self.ids = dns.rdata.Rdata._as_tuple( - ids, lambda x: dns.rdata.Rdata._as_bytes(x, True, 255, False)) + ids, lambda x: dns.rdata.Rdata._as_bytes(x, True, 255, False) + ) @classmethod def from_value(cls, value): return cls(_split(_unescape(value))) def to_text(self): - value = ','.join([_escapify(id) for id in self.ids]) + value = ",".join([_escapify(id) for id in self.ids]) return '"' + dns.rdata._escapify(value.encode()) + '"' @classmethod @@ -253,7 +258,7 @@ class ALPNParam(Param): def to_wire(self, file, origin=None): # pylint: disable=W0613 for id in self.ids: - file.write(struct.pack('!B', len(id))) + file.write(struct.pack("!B", len(id))) file.write(id) @@ -269,10 +274,10 @@ class NoDefaultALPNParam(Param): @classmethod def from_value(cls, value): - if value is None or value == '': + if value is None or value == "": return None else: - raise ValueError('no-default-alpn with non-empty value') + raise ValueError("no-default-alpn with non-empty value") def to_text(self): raise NotImplementedError # pragma: no cover @@ -306,22 +311,23 @@ class PortParam(Param): return cls(port) def to_wire(self, file, origin=None): # pylint: disable=W0613 - file.write(struct.pack('!H', self.port)) + file.write(struct.pack("!H", self.port)) @dns.immutable.immutable class IPv4HintParam(Param): def __init__(self, addresses): self.addresses = dns.rdata.Rdata._as_tuple( - addresses, dns.rdata.Rdata._as_ipv4_address) + addresses, dns.rdata.Rdata._as_ipv4_address + ) @classmethod def from_value(cls, value): - addresses = value.split(',') + addresses = value.split(",") return cls(addresses) def to_text(self): - return '"' + ','.join(self.addresses) + '"' + return '"' + ",".join(self.addresses) + '"' @classmethod def from_wire_parser(cls, parser, origin=None): # pylint: disable=W0613 @@ -340,15 +346,16 @@ class IPv4HintParam(Param): class IPv6HintParam(Param): def __init__(self, addresses): self.addresses = dns.rdata.Rdata._as_tuple( - addresses, dns.rdata.Rdata._as_ipv6_address) + addresses, dns.rdata.Rdata._as_ipv6_address + ) @classmethod def from_value(cls, value): - addresses = value.split(',') + addresses = value.split(",") return cls(addresses) def to_text(self): - return '"' + ','.join(self.addresses) + '"' + return '"' + ",".join(self.addresses) + '"' @classmethod def from_wire_parser(cls, parser, origin=None): # pylint: disable=W0613 @@ -370,13 +377,13 @@ class ECHParam(Param): @classmethod def from_value(cls, value): - if '\\' in value: - raise ValueError('escape in ECH value') + if "\\" in value: + raise ValueError("escape in ECH value") value = base64.b64decode(value.encode()) return cls(value) def to_text(self): - b64 = base64.b64encode(self.ech).decode('ascii') + b64 = base64.b64encode(self.ech).decode("ascii") return f'"{b64}"' @classmethod @@ -407,7 +414,7 @@ def _validate_and_define(params, key, value): emptiness = cls.emptiness() if value is None: if emptiness == Emptiness.NEVER: - raise SyntaxError('value cannot be empty') + raise SyntaxError("value cannot be empty") value = cls.from_value(value) else: if force_generic: @@ -424,7 +431,7 @@ class SVCBBase(dns.rdata.Rdata): # see: draft-ietf-dnsop-svcb-https-01 - __slots__ = ['priority', 'target', 'params'] + __slots__ = ["priority", "target", "params"] def __init__(self, rdclass, rdtype, priority, target, params): super().__init__(rdclass, rdtype) @@ -443,12 +450,13 @@ class SVCBBase(dns.rdata.Rdata): # Note we have to say "not in" as we have None as a value # so a get() and a not None test would be wrong. if key not in params: - raise ValueError(f'key {key:d} declared mandatory but not ' - 'present') + raise ValueError( + f"key {key:d} declared mandatory but not " "present" + ) # The no-default-alpn parameter requires the alpn parameter. if ParamKey.NO_DEFAULT_ALPN in params: if ParamKey.ALPN not in params: - raise ValueError('no-default-alpn present, but alpn missing') + raise ValueError("no-default-alpn present, but alpn missing") def to_text(self, origin=None, relativize=True, **kw): target = self.target.choose_relativity(origin, relativize) @@ -458,23 +466,24 @@ class SVCBBase(dns.rdata.Rdata): if value is None: params.append(key_to_text(key)) else: - kv = key_to_text(key) + '=' + value.to_text() + kv = key_to_text(key) + "=" + value.to_text() params.append(kv) if len(params) > 0: - space = ' ' + space = " " else: - space = '' - return '%d %s%s%s' % (self.priority, target, space, ' '.join(params)) + space = "" + return "%d %s%s%s" % (self.priority, target, space, " ".join(params)) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): priority = tok.get_uint16() target = tok.get_name(origin, relativize, relativize_to) if priority == 0: token = tok.get() if not token.is_eol_or_eof(): - raise SyntaxError('parameters in AliasMode') + raise SyntaxError("parameters in AliasMode") tok.unget(token) params = {} while True: @@ -483,20 +492,20 @@ class SVCBBase(dns.rdata.Rdata): tok.unget(token) break if token.ttype != dns.tokenizer.IDENTIFIER: - raise SyntaxError('parameter is not an identifier') - equals = token.value.find('=') + raise SyntaxError("parameter is not an identifier") + equals = token.value.find("=") if equals == len(token.value) - 1: # 'key=', so next token should be a quoted string without # any intervening whitespace. key = token.value[:-1] token = tok.get(want_leading=True) if token.ttype != dns.tokenizer.QUOTED_STRING: - raise SyntaxError('whitespace after =') + raise SyntaxError("whitespace after =") value = token.value elif equals > 0: # key=value key = token.value[:equals] - value = token.value[equals + 1:] + value = token.value[equals + 1 :] elif equals == 0: # =key raise SyntaxError('parameter cannot start with "="') @@ -532,13 +541,13 @@ class SVCBBase(dns.rdata.Rdata): priority = parser.get_uint16() target = parser.get_name(origin) if priority == 0 and parser.remaining() != 0: - raise dns.exception.FormError('parameters in AliasMode') + raise dns.exception.FormError("parameters in AliasMode") params = {} prior_key = -1 while parser.remaining() > 0: key = parser.get_uint16() if key < prior_key: - raise dns.exception.FormError('keys not in order') + raise dns.exception.FormError("keys not in order") prior_key = key vlen = parser.get_uint16() pcls = _class_for_key.get(key, GenericParam) diff --git a/dns/rdtypes/tlsabase.py b/dns/rdtypes/tlsabase.py index 786fca5..a3fdc35 100644 --- a/dns/rdtypes/tlsabase.py +++ b/dns/rdtypes/tlsabase.py @@ -30,10 +30,9 @@ class TLSABase(dns.rdata.Rdata): # see: RFC 6698 - __slots__ = ['usage', 'selector', 'mtype', 'cert'] + __slots__ = ["usage", "selector", "mtype", "cert"] - def __init__(self, rdclass, rdtype, usage, selector, - mtype, cert): + def __init__(self, rdclass, rdtype, usage, selector, mtype, cert): super().__init__(rdclass, rdtype) self.usage = self._as_uint8(usage) self.selector = self._as_uint8(selector) @@ -42,17 +41,18 @@ class TLSABase(dns.rdata.Rdata): def to_text(self, origin=None, relativize=True, **kw): kw = kw.copy() - chunksize = kw.pop('chunksize', 128) - return '%d %d %d %s' % (self.usage, - self.selector, - self.mtype, - dns.rdata._hexify(self.cert, - chunksize=chunksize, - **kw)) + chunksize = kw.pop("chunksize", 128) + return "%d %d %d %s" % ( + self.usage, + self.selector, + self.mtype, + dns.rdata._hexify(self.cert, chunksize=chunksize, **kw), + ) @classmethod - def from_text(cls, rdclass, rdtype, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, rdclass, rdtype, tok, origin=None, relativize=True, relativize_to=None + ): usage = tok.get_uint8() selector = tok.get_uint8() mtype = tok.get_uint8() diff --git a/dns/rdtypes/txtbase.py b/dns/rdtypes/txtbase.py index afef98e..d4cb9bb 100644 --- a/dns/rdtypes/txtbase.py +++ b/dns/rdtypes/txtbase.py @@ -32,11 +32,14 @@ class TXTBase(dns.rdata.Rdata): """Base class for rdata that is like a TXT record (see RFC 1035).""" - __slots__ = ['strings'] - - def __init__(self, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - strings: Iterable[Union[bytes, str]]): + __slots__ = ["strings"] + + def __init__( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + strings: Iterable[Union[bytes, str]], + ): """Initialize a TXT-like rdata. *rdclass*, an ``int`` is the rdataclass of the Rdata. @@ -46,28 +49,41 @@ class TXTBase(dns.rdata.Rdata): *strings*, a tuple of ``bytes`` """ super().__init__(rdclass, rdtype) - self.strings: Tuple[bytes] = self._as_tuple(strings, lambda x: self._as_bytes(x, True, 255)) - - def to_text(self, origin: Optional[dns.name.Name]=None, relativize: bool=True, **kw: Dict[str, Any]) -> str: - txt = '' - prefix = '' + self.strings: Tuple[bytes] = self._as_tuple( + strings, lambda x: self._as_bytes(x, True, 255) + ) + + def to_text( + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + **kw: Dict[str, Any] + ) -> str: + txt = "" + prefix = "" for s in self.strings: txt += '{}"{}"'.format(prefix, dns.rdata._escapify(s)) - prefix = ' ' + prefix = " " return txt @classmethod - def from_text(cls, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - tok: dns.tokenizer.Tokenizer, origin: Optional[dns.name.Name]=None, - relativize: bool=True, relativize_to: Optional[dns.name.Name]=None) -> dns.rdata.Rdata: + def from_text( + cls, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + tok: dns.tokenizer.Tokenizer, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, + ) -> dns.rdata.Rdata: strings = [] for token in tok.get_remaining(): token = token.unescape_to_bytes() # The 'if' below is always true in the current code, but we # are leaving this check in in case things change some day. - if not (token.is_quoted_string() or - token.is_identifier()): # pragma: no cover + if not ( + token.is_quoted_string() or token.is_identifier() + ): # pragma: no cover raise dns.exception.SyntaxError("expected a string") if len(token.value) > 255: raise dns.exception.SyntaxError("string too long") @@ -80,7 +96,7 @@ class TXTBase(dns.rdata.Rdata): for s in self.strings: l = len(s) assert l < 256 - file.write(struct.pack('!B', l)) + file.write(struct.pack("!B", l)) file.write(s) @classmethod diff --git a/dns/rdtypes/util.py b/dns/rdtypes/util.py index 9bf8f7e..74596f0 100644 --- a/dns/rdtypes/util.py +++ b/dns/rdtypes/util.py @@ -28,6 +28,7 @@ import dns.rdata class Gateway: """A helper class for the IPSECKEY gateway and AMTRELAY relay fields""" + name = "" def __init__(self, type, gateway=None): @@ -67,15 +68,17 @@ class Gateway: raise ValueError(self._invalid_type(self.type)) # pragma: no cover @classmethod - def from_text(cls, gateway_type, tok, origin=None, relativize=True, - relativize_to=None): + def from_text( + cls, gateway_type, tok, origin=None, relativize=True, relativize_to=None + ): if gateway_type in (0, 1, 2): gateway = tok.get_string() elif gateway_type == 3: gateway = tok.get_name(origin, relativize, relativize_to) else: raise dns.exception.SyntaxError( - cls._invalid_type(gateway_type)) # pragma: no cover + cls._invalid_type(gateway_type) + ) # pragma: no cover return cls(gateway_type, gateway) # pylint: disable=unused-argument @@ -90,6 +93,7 @@ class Gateway: self.gateway.to_wire(file, None, origin, False) else: raise ValueError(self._invalid_type(self.type)) # pragma: no cover + # pylint: enable=unused-argument @classmethod @@ -109,6 +113,7 @@ class Gateway: class Bitmap: """A helper class for the NSEC/NSEC3/CSYNC type bitmaps""" + type_name = "" def __init__(self, windows=None): @@ -136,7 +141,7 @@ class Bitmap: if byte & (0x80 >> j): rdtype = window * 256 + i * 8 + j bits.append(dns.rdatatype.to_text(rdtype)) - text += (' ' + ' '.join(bits)) + text += " " + " ".join(bits) return text @classmethod @@ -151,7 +156,7 @@ class Bitmap: window = 0 octets = 0 prior_rdtype = 0 - bitmap = bytearray(b'\0' * 32) + bitmap = bytearray(b"\0" * 32) windows = [] for rdtype in rdtypes: if rdtype == prior_rdtype: @@ -161,7 +166,7 @@ class Bitmap: if new_window != window: if octets != 0: windows.append((window, bytes(bitmap[0:octets]))) - bitmap = bytearray(b'\0' * 32) + bitmap = bytearray(b"\0" * 32) window = new_window offset = rdtype % 256 byte = offset // 8 @@ -174,7 +179,7 @@ class Bitmap: def to_wire(self, file): for (window, bitmap) in self.windows: - file.write(struct.pack('!BB', window, len(bitmap))) + file.write(struct.pack("!BB", window, len(bitmap))) file.write(bitmap) @classmethod @@ -193,6 +198,7 @@ def _priority_table(items): by_priority[rdata._processing_priority()].append(rdata) return by_priority + def priority_processing_order(iterable): items = list(iterable) if len(items) == 1: @@ -205,8 +211,10 @@ def priority_processing_order(iterable): ordered.extend(rdatas) return ordered + _no_weight = 0.1 + def weighted_processing_order(iterable): items = list(iterable) if len(items) == 1: @@ -215,8 +223,7 @@ def weighted_processing_order(iterable): ordered = [] for k in sorted(by_priority.keys()): rdatas = by_priority[k] - total = sum(rdata._processing_weight() or _no_weight - for rdata in rdatas) + total = sum(rdata._processing_weight() or _no_weight for rdata in rdatas) while len(rdatas) > 1: r = random.uniform(0, total) for (n, rdata) in enumerate(rdatas): @@ -230,15 +237,16 @@ def weighted_processing_order(iterable): ordered.append(rdatas[0]) return ordered + def parse_formatted_hex(formatted, num_chunks, chunk_size, separator): if len(formatted) != num_chunks * (chunk_size + 1) - 1: - raise ValueError('invalid formatted hex string') - value = b'' + raise ValueError("invalid formatted hex string") + value = b"" for _ in range(num_chunks): chunk = formatted[0:chunk_size] - value += int(chunk, 16).to_bytes(chunk_size // 2, 'big') + value += int(chunk, 16).to_bytes(chunk_size // 2, "big") formatted = formatted[chunk_size:] if len(formatted) > 0 and formatted[0] != separator: - raise ValueError('invalid formatted hex string') + raise ValueError("invalid formatted hex string") formatted = formatted[1:] return value diff --git a/dns/renderer.py b/dns/renderer.py index 4e4391c..95e8bd3 100644 --- a/dns/renderer.py +++ b/dns/renderer.py @@ -88,8 +88,8 @@ class Renderer: self.compress = {} self.section = QUESTION self.counts = [0, 0, 0, 0] - self.output.write(b'\x00' * 12) - self.mac = '' + self.output.write(b"\x00" * 12) + self.mac = "" def _rollback(self, where): """Truncate the output buffer at offset *where*, and remove any @@ -160,8 +160,7 @@ class Renderer: self._set_section(section) with self._track_size(): - n = rdataset.to_wire(name, self.output, self.compress, self.origin, - **kw) + n = rdataset.to_wire(name, self.output, self.compress, self.origin, **kw) self.counts[section] += n def add_edns(self, edns, ednsflags, payload, options=None): @@ -169,12 +168,21 @@ class Renderer: # make sure the EDNS version in ednsflags agrees with edns ednsflags &= 0xFF00FFFF - ednsflags |= (edns << 16) + ednsflags |= edns << 16 opt = dns.message.Message._make_opt(ednsflags, payload, options) self.add_rrset(ADDITIONAL, opt) - def add_tsig(self, keyname, secret, fudge, id, tsig_error, other_data, - request_mac, algorithm=dns.tsig.default_algorithm): + def add_tsig( + self, + keyname, + secret, + fudge, + id, + tsig_error, + other_data, + request_mac, + algorithm=dns.tsig.default_algorithm, + ): """Add a TSIG signature to the message.""" s = self.output.getvalue() @@ -183,15 +191,24 @@ class Renderer: key = secret else: key = dns.tsig.Key(keyname, secret, algorithm) - tsig = dns.message.Message._make_tsig(keyname, algorithm, 0, fudge, - b'', id, tsig_error, other_data) - (tsig, _) = dns.tsig.sign(s, key, tsig[0], int(time.time()), - request_mac) + tsig = dns.message.Message._make_tsig( + keyname, algorithm, 0, fudge, b"", id, tsig_error, other_data + ) + (tsig, _) = dns.tsig.sign(s, key, tsig[0], int(time.time()), request_mac) self._write_tsig(tsig, keyname) - def add_multi_tsig(self, ctx, keyname, secret, fudge, id, tsig_error, - other_data, request_mac, - algorithm=dns.tsig.default_algorithm): + def add_multi_tsig( + self, + ctx, + keyname, + secret, + fudge, + id, + tsig_error, + other_data, + request_mac, + algorithm=dns.tsig.default_algorithm, + ): """Add a TSIG signature to the message. Unlike add_tsig(), this can be used for a series of consecutive DNS envelopes, e.g. for a zone transfer over TCP [RFC2845, 4.4]. @@ -206,10 +223,12 @@ class Renderer: key = secret else: key = dns.tsig.Key(keyname, secret, algorithm) - tsig = dns.message.Message._make_tsig(keyname, algorithm, 0, fudge, - b'', id, tsig_error, other_data) - (tsig, ctx) = dns.tsig.sign(s, key, tsig[0], int(time.time()), - request_mac, ctx, True) + tsig = dns.message.Message._make_tsig( + keyname, algorithm, 0, fudge, b"", id, tsig_error, other_data + ) + (tsig, ctx) = dns.tsig.sign( + s, key, tsig[0], int(time.time()), request_mac, ctx, True + ) self._write_tsig(tsig, keyname) return ctx @@ -217,17 +236,18 @@ class Renderer: self._set_section(ADDITIONAL) with self._track_size(): keyname.to_wire(self.output, self.compress, self.origin) - self.output.write(struct.pack('!HHIH', dns.rdatatype.TSIG, - dns.rdataclass.ANY, 0, 0)) + self.output.write( + struct.pack("!HHIH", dns.rdatatype.TSIG, dns.rdataclass.ANY, 0, 0) + ) rdata_start = self.output.tell() tsig.to_wire(self.output) after = self.output.tell() self.output.seek(rdata_start - 2) - self.output.write(struct.pack('!H', after - rdata_start)) + self.output.write(struct.pack("!H", after - rdata_start)) self.counts[ADDITIONAL] += 1 self.output.seek(10) - self.output.write(struct.pack('!H', self.counts[ADDITIONAL])) + self.output.write(struct.pack("!H", self.counts[ADDITIONAL])) self.output.seek(0, io.SEEK_END) def write_header(self): @@ -239,9 +259,17 @@ class Renderer: """ self.output.seek(0) - self.output.write(struct.pack('!HHHHHH', self.id, self.flags, - self.counts[0], self.counts[1], - self.counts[2], self.counts[3])) + self.output.write( + struct.pack( + "!HHHHHH", + self.id, + self.flags, + self.counts[0], + self.counts[1], + self.counts[2], + self.counts[3], + ) + ) self.output.seek(0, io.SEEK_END) def get_wire(self): diff --git a/dns/resolver.py b/dns/resolver.py index 0b13253..f930369 100644 --- a/dns/resolver.py +++ b/dns/resolver.py @@ -26,10 +26,11 @@ import sys import time import random import warnings + try: import threading as _threading except ImportError: # pragma: no cover - import dummy_threading as _threading # type: ignore + import dummy_threading as _threading # type: ignore import dns.exception import dns.edns @@ -46,22 +47,24 @@ import dns.rdatatype import dns.reversename import dns.tsig -if sys.platform == 'win32': +if sys.platform == "win32": import dns.win32util + class NXDOMAIN(dns.exception.DNSException): """The DNS query name does not exist.""" - supp_kwargs = {'qnames', 'responses'} + + supp_kwargs = {"qnames", "responses"} fmt = None # we have our own __str__ implementation # pylint: disable=arguments-differ - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - def _check_kwargs(self, qnames, - responses=None): + def _check_kwargs(self, qnames, responses=None): if not isinstance(qnames, (list, tuple, set)): raise AttributeError("qnames must be a list, tuple or set") if len(qnames) == 0: @@ -74,23 +77,23 @@ class NXDOMAIN(dns.exception.DNSException): return kwargs def __str__(self): - if 'qnames' not in self.kwargs: + if "qnames" not in self.kwargs: return super().__str__() - qnames = self.kwargs['qnames'] + qnames = self.kwargs["qnames"] if len(qnames) > 1: - msg = 'None of DNS query names exist' + msg = "None of DNS query names exist" else: - msg = 'The DNS query name does not exist' - qnames = ', '.join(map(str, qnames)) + msg = "The DNS query name does not exist" + qnames = ", ".join(map(str, qnames)) return "{}: {}".format(msg, qnames) @property def canonical_name(self): """Return the unresolved canonical name.""" - if 'qnames' not in self.kwargs: + if "qnames" not in self.kwargs: raise TypeError("parametrized exception required") - for qname in self.kwargs['qnames']: - response = self.kwargs['responses'][qname] + for qname in self.kwargs["qnames"]: + response = self.kwargs["responses"][qname] try: cname = response.canonical_name() if cname != qname: @@ -99,14 +102,14 @@ class NXDOMAIN(dns.exception.DNSException): # We can just eat this exception as it means there was # something wrong with the response. pass - return self.kwargs['qnames'][0] + return self.kwargs["qnames"][0] def __add__(self, e_nx): """Augment by results from another NXDOMAIN exception.""" - qnames0 = list(self.kwargs.get('qnames', [])) - responses0 = dict(self.kwargs.get('responses', {})) - responses1 = e_nx.kwargs.get('responses', {}) - for qname1 in e_nx.kwargs.get('qnames', []): + qnames0 = list(self.kwargs.get("qnames", [])) + responses0 = dict(self.kwargs.get("responses", {})) + responses1 = e_nx.kwargs.get("responses", {}) + for qname1 in e_nx.kwargs.get("qnames", []): if qname1 not in qnames0: qnames0.append(qname1) if qname1 in responses1: @@ -118,7 +121,7 @@ class NXDOMAIN(dns.exception.DNSException): Returns a list of ``dns.name.Name``. """ - return self.kwargs['qnames'] + return self.kwargs["qnames"] def responses(self): """A map from queried names to their NXDOMAIN responses. @@ -126,29 +129,34 @@ class NXDOMAIN(dns.exception.DNSException): Returns a dict mapping a ``dns.name.Name`` to a ``dns.message.Message``. """ - return self.kwargs['responses'] + return self.kwargs["responses"] def response(self, qname): """The response for query *qname*. Returns a ``dns.message.Message``. """ - return self.kwargs['responses'][qname] + return self.kwargs["responses"][qname] class YXDOMAIN(dns.exception.DNSException): """The DNS query name is too long after DNAME substitution.""" -ErrorTuple = Tuple[Optional[str], bool, int, Union[Exception, str], Optional[dns.message.Message]] +ErrorTuple = Tuple[ + Optional[str], bool, int, Union[Exception, str], Optional[dns.message.Message] +] def _errors_to_text(errors: List[ErrorTuple]) -> List[str]: """Turn a resolution errors trace into a list of text.""" texts = [] for err in errors: - texts.append('Server {} {} port {} answered {}'.format(err[0], - 'TCP' if err[1] else 'UDP', err[2], err[3])) + texts.append( + "Server {} {} port {} answered {}".format( + err[0], "TCP" if err[1] else "UDP", err[2], err[3] + ) + ) return texts @@ -157,16 +165,18 @@ class LifetimeTimeout(dns.exception.Timeout): msg = "The resolution lifetime expired." fmt = "%s after {timeout:.3f} seconds: {errors}" % msg[:-1] - supp_kwargs = {'timeout', 'errors'} + supp_kwargs = {"timeout", "errors"} - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def _fmt_kwargs(self, **kwargs): - srv_msgs = _errors_to_text(kwargs['errors']) - return super()._fmt_kwargs(timeout=kwargs['timeout'], - errors='; '.join(srv_msgs)) + srv_msgs = _errors_to_text(kwargs["errors"]) + return super()._fmt_kwargs( + timeout=kwargs["timeout"], errors="; ".join(srv_msgs) + ) # We added more detail to resolution timeouts, but they are still @@ -177,19 +187,20 @@ Timeout = LifetimeTimeout class NoAnswer(dns.exception.DNSException): """The DNS response does not contain an answer to the question.""" - fmt = 'The DNS response does not contain an answer ' + \ - 'to the question: {query}' - supp_kwargs = {'response'} - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + fmt = "The DNS response does not contain an answer " + "to the question: {query}" + supp_kwargs = {"response"} + + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def _fmt_kwargs(self, **kwargs): - return super()._fmt_kwargs(query=kwargs['response'].question) + return super()._fmt_kwargs(query=kwargs["response"].question) def response(self): - return self.kwargs['response'] + return self.kwargs["response"] class NoNameservers(dns.exception.DNSException): @@ -203,16 +214,18 @@ class NoNameservers(dns.exception.DNSException): msg = "All nameservers failed to answer the query." fmt = "%s {query}: {errors}" % msg[:-1] - supp_kwargs = {'request', 'errors'} + supp_kwargs = {"request", "errors"} - # We do this as otherwise mypy complains about unexpected keyword argument idna_exception + # We do this as otherwise mypy complains about unexpected keyword argument + # idna_exception def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def _fmt_kwargs(self, **kwargs): - srv_msgs = _errors_to_text(kwargs['errors']) - return super()._fmt_kwargs(query=kwargs['request'].question, - errors='; '.join(srv_msgs)) + srv_msgs = _errors_to_text(kwargs["errors"]) + return super()._fmt_kwargs( + query=kwargs["request"].question, errors="; ".join(srv_msgs) + ) class NotAbsolute(dns.exception.DNSException): @@ -226,9 +239,11 @@ class NoRootSOA(dns.exception.DNSException): class NoMetaqueries(dns.exception.DNSException): """DNS metaqueries are not allowed.""" + class NoResolverConfiguration(dns.exception.DNSException): """Resolver configuration could not be read or specified no nameservers.""" + class Answer: """DNS stub resolver answer. @@ -245,9 +260,15 @@ class Answer: RRset's name might not be the query name. """ - def __init__(self, qname: dns.name.Name, rdtype: dns.rdatatype.RdataType, - rdclass: dns.rdataclass.RdataClass, response: dns.message.QueryMessage, - nameserver: Optional[str]=None, port: Optional[int]=None): + def __init__( + self, + qname: dns.name.Name, + rdtype: dns.rdatatype.RdataType, + rdclass: dns.rdataclass.RdataClass, + response: dns.message.QueryMessage, + nameserver: Optional[str] = None, + port: Optional[int] = None, + ): self.qname = qname self.rdtype = rdtype self.rdclass = rdclass @@ -262,15 +283,15 @@ class Answer: self.expiration = time.time() + self.chaining_result.minimum_ttl def __getattr__(self, attr): # pragma: no cover - if attr == 'name': + if attr == "name": return self.rrset.name - elif attr == 'ttl': + elif attr == "ttl": return self.rrset.ttl - elif attr == 'covers': + elif attr == "covers": return self.rrset.covers - elif attr == 'rdclass': + elif attr == "rdclass": return self.rrset.rdclass - elif attr == 'rdtype': + elif attr == "rdtype": return self.rrset.rdtype else: raise AttributeError(attr) @@ -293,8 +314,7 @@ class Answer: class CacheStatistics: - """Cache Statistics - """ + """Cache Statistics""" def __init__(self, hits=0, misses=0): self.hits = hits @@ -304,7 +324,7 @@ class CacheStatistics: self.hits = 0 self.misses = 0 - def clone(self) -> 'CacheStatistics': + def clone(self) -> "CacheStatistics": return CacheStatistics(self.hits, self.misses) @@ -345,7 +365,7 @@ CacheKey = Tuple[dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataCla class Cache(CacheBase): """Simple thread-safe DNS answer cache.""" - def __init__(self, cleaning_interval: float=300.0): + def __init__(self, cleaning_interval: float = 300.0): """*cleaning_interval*, a ``float`` is the number of seconds between periodic cleanings. """ @@ -374,8 +394,8 @@ class Cache(CacheBase): Returns None if no answer is cached for the key. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. Returns a ``dns.resolver.Answer`` or ``None``. """ @@ -392,8 +412,8 @@ class Cache(CacheBase): def put(self, key: CacheKey, value: Answer) -> None: """Associate key and value in the cache. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. *value*, a ``dns.resolver.Answer``, the answer. """ @@ -402,14 +422,14 @@ class Cache(CacheBase): self._maybe_clean() self.data[key] = value - def flush(self, key: Optional[CacheKey]=None) -> None: + def flush(self, key: Optional[CacheKey] = None) -> None: """Flush the cache. - If *key* is not ``None``, only that item is flushed. Otherwise - the entire cache is flushed. + If *key* is not ``None``, only that item is flushed. Otherwise the entire cache + is flushed. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. """ with self.lock: @@ -452,7 +472,7 @@ class LRUCache(CacheBase): for a new one. """ - def __init__(self, max_size: int=100000): + def __init__(self, max_size: int = 100000): """*max_size*, an ``int``, is the maximum number of nodes to cache; it must be greater than 0. """ @@ -474,8 +494,8 @@ class LRUCache(CacheBase): Returns None if no answer is cached for the key. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. Returns a ``dns.resolver.Answer`` or ``None``. """ @@ -509,8 +529,8 @@ class LRUCache(CacheBase): def put(self, key: CacheKey, value: Answer) -> None: """Associate key and value in the cache. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. *value*, a ``dns.resolver.Answer``, the answer. """ @@ -528,14 +548,14 @@ class LRUCache(CacheBase): node.link_after(self.sentinel) self.data[key] = node - def flush(self, key: Optional[CacheKey]=None) -> None: + def flush(self, key: Optional[CacheKey] = None) -> None: """Flush the cache. - If *key* is not ``None``, only that item is flushed. Otherwise - the entire cache is flushed. + If *key* is not ``None``, only that item is flushed. Otherwise the entire cache + is flushed. - *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` tuple whose values are the - query name, rdtype, and rdclass respectively. + *key*, a ``(dns.name.Name, dns.rdatatype.RdataType, dns.rdataclass.RdataClass)`` + tuple whose values are the query name, rdtype, and rdclass respectively. """ with self.lock: @@ -552,6 +572,7 @@ class LRUCache(CacheBase): gnode = next self.data = {} + class _Resolution: """Helper class for dns.resolver.Resolver.resolve(). @@ -564,10 +585,16 @@ class _Resolution: resolver data structures directly. """ - def __init__(self, resolver: 'BaseResolver', qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - rdclass: Union[dns.rdataclass.RdataClass, str], - tcp: bool, raise_on_no_answer: bool, search: Optional[bool]): + def __init__( + self, + resolver: "BaseResolver", + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + rdclass: Union[dns.rdataclass.RdataClass, str], + tcp: bool, + raise_on_no_answer: bool, + search: Optional[bool], + ): if isinstance(qname, str): qname = dns.name.from_text(qname, None) the_rdtype = dns.rdatatype.RdataType.make(rdtype) @@ -596,7 +623,9 @@ class _Resolution: self.request: Optional[dns.message.QueryMessage] = None self.backoff = 0.0 - def next_request(self) -> Tuple[Optional[dns.message.QueryMessage], Optional[Answer]]: + def next_request( + self, + ) -> Tuple[Optional[dns.message.QueryMessage], Optional[Answer]]: """Get the next request to send, and check the cache. Returns a (request, answer) tuple. At most one of request or @@ -611,32 +640,37 @@ class _Resolution: # Do we know the answer? if self.resolver.cache: - answer = self.resolver.cache.get((self.qname, self.rdtype, - self.rdclass)) + answer = self.resolver.cache.get( + (self.qname, self.rdtype, self.rdclass) + ) if answer is not None: if answer.rrset is None and self.raise_on_no_answer: raise NoAnswer(response=answer.response) else: return (None, answer) - answer = self.resolver.cache.get((self.qname, - dns.rdatatype.ANY, - self.rdclass)) - if answer is not None and \ - answer.response.rcode() == dns.rcode.NXDOMAIN: + answer = self.resolver.cache.get( + (self.qname, dns.rdatatype.ANY, self.rdclass) + ) + if answer is not None and answer.response.rcode() == dns.rcode.NXDOMAIN: # cached NXDOMAIN; record it and continue to next # name. self.nxdomain_responses[self.qname] = answer.response continue # Build the request - request = dns.message.make_query(self.qname, self.rdtype, - self.rdclass) + request = dns.message.make_query(self.qname, self.rdtype, self.rdclass) if self.resolver.keyname is not None: - request.use_tsig(self.resolver.keyring, self.resolver.keyname, - algorithm=self.resolver.keyalgorithm) - request.use_edns(self.resolver.edns, self.resolver.ednsflags, - self.resolver.payload, - options=self.resolver.ednsoptions) + request.use_tsig( + self.resolver.keyring, + self.resolver.keyname, + algorithm=self.resolver.keyalgorithm, + ) + request.use_edns( + self.resolver.edns, + self.resolver.ednsflags, + self.resolver.payload, + options=self.resolver.ednsoptions, + ) if self.resolver.flags is not None: request.flags = self.resolver.flags @@ -658,8 +692,7 @@ class _Resolution: # it's only NXDOMAINs as anything else would have returned # before now.) # - raise NXDOMAIN(qnames=self.qnames_to_try, - responses=self.nxdomain_responses) + raise NXDOMAIN(qnames=self.qnames_to_try, responses=self.nxdomain_responses) def next_nameserver(self) -> Tuple[str, int, bool, float]: if self.retry_with_tcp: @@ -678,13 +711,15 @@ class _Resolution: self.backoff = min(self.backoff * 2, 2) self.nameserver = self.current_nameservers.pop(0) - self.port = self.resolver.nameserver_ports.get(self.nameserver, - self.resolver.port) + self.port = self.resolver.nameserver_ports.get( + self.nameserver, self.resolver.port + ) self.tcp_attempt = self.tcp return (self.nameserver, self.port, self.tcp_attempt, backoff) - def query_result(self, response: Optional[dns.message.Message], - ex: Optional[Exception]) -> Tuple[Optional[Answer], bool]: + def query_result( + self, response: Optional[dns.message.Message], ex: Optional[Exception] + ) -> Tuple[Optional[Answer], bool]: # # returns an (answer: Answer, end_loop: bool) tuple. # @@ -692,12 +727,15 @@ class _Resolution: if ex: # Exception during I/O or from_wire() assert response is None - self.errors.append((self.nameserver, self.tcp_attempt, self.port, - ex, response)) - if isinstance(ex, dns.exception.FormError) or \ - isinstance(ex, EOFError) or \ - isinstance(ex, OSError) or \ - isinstance(ex, NotImplementedError): + self.errors.append( + (self.nameserver, self.tcp_attempt, self.port, ex, response) + ) + if ( + isinstance(ex, dns.exception.FormError) + or isinstance(ex, EOFError) + or isinstance(ex, OSError) + or isinstance(ex, NotImplementedError) + ): # This nameserver is no good, take it out of the mix. self.nameservers.remove(self.nameserver) elif isinstance(ex, dns.message.Truncated): @@ -713,17 +751,23 @@ class _Resolution: rcode = response.rcode() if rcode == dns.rcode.NOERROR: try: - answer = Answer(self.qname, self.rdtype, self.rdclass, response, - self.nameserver, self.port) + answer = Answer( + self.qname, + self.rdtype, + self.rdclass, + response, + self.nameserver, + self.port, + ) except Exception as e: - self.errors.append((self.nameserver, self.tcp_attempt, - self.port, e, response)) + self.errors.append( + (self.nameserver, self.tcp_attempt, self.port, e, response) + ) # The nameserver is no good, take it out of the mix. self.nameservers.remove(self.nameserver) return (None, False) if self.resolver.cache: - self.resolver.cache.put((self.qname, self.rdtype, - self.rdclass), answer) + self.resolver.cache.put((self.qname, self.rdtype, self.rdclass), answer) if answer.rrset is None and self.raise_on_no_answer: raise NoAnswer(response=answer.response) return (answer, True) @@ -731,26 +775,29 @@ class _Resolution: # Further validate the response by making an Answer, even # if we aren't going to cache it. try: - answer = Answer(self.qname, dns.rdatatype.ANY, - dns.rdataclass.IN, response) + answer = Answer( + self.qname, dns.rdatatype.ANY, dns.rdataclass.IN, response + ) except Exception as e: - self.errors.append((self.nameserver, self.tcp_attempt, - self.port, e, response)) + self.errors.append( + (self.nameserver, self.tcp_attempt, self.port, e, response) + ) # The nameserver is no good, take it out of the mix. self.nameservers.remove(self.nameserver) return (None, False) self.nxdomain_responses[self.qname] = response if self.resolver.cache: - self.resolver.cache.put((self.qname, - dns.rdatatype.ANY, - self.rdclass), answer) + self.resolver.cache.put( + (self.qname, dns.rdatatype.ANY, self.rdclass), answer + ) # Make next_nameserver() return None, so caller breaks its # inner loop and calls next_request(). return (None, True) elif rcode == dns.rcode.YXDOMAIN: yex = YXDOMAIN() - self.errors.append((self.nameserver, self.tcp_attempt, - self.port, yex, response)) + self.errors.append( + (self.nameserver, self.tcp_attempt, self.port, yex, response) + ) raise yex else: # @@ -759,8 +806,15 @@ class _Resolution: # if rcode != dns.rcode.SERVFAIL or not self.resolver.retry_servfail: self.nameservers.remove(self.nameserver) - self.errors.append((self.nameserver, self.tcp_attempt, self.port, - dns.rcode.to_text(rcode), response)) + self.errors.append( + ( + self.nameserver, + self.tcp_attempt, + self.port, + dns.rcode.to_text(rcode), + response, + ) + ) return (None, False) @@ -791,7 +845,7 @@ class BaseResolver: rotate: bool ndots: Optional[int] - def __init__(self, filename: str='/etc/resolv.conf', configure: bool=True): + def __init__(self, filename: str = "/etc/resolv.conf", configure: bool = True): """*filename*, a ``str`` or file object, specifying a file in standard /etc/resolv.conf format. This parameter is meaningful only when *configure* is true and the platform is POSIX. @@ -805,7 +859,7 @@ class BaseResolver: self.reset() if configure: - if sys.platform == 'win32': + if sys.platform == "win32": self.read_registry() elif filename: self.read_resolv_conf(filename) @@ -859,10 +913,10 @@ class BaseResolver: f = stack.enter_context(open(f)) except OSError: # /etc/resolv.conf doesn't exist, can't be read, etc. - raise NoResolverConfiguration(f'cannot open {f}') + raise NoResolverConfiguration(f"cannot open {f}") for l in f: - if len(l) == 0 or l[0] == '#' or l[0] == ';': + if len(l) == 0 or l[0] == "#" or l[0] == ";": continue tokens = l.split() @@ -870,37 +924,37 @@ class BaseResolver: if len(tokens) < 2: continue - if tokens[0] == 'nameserver': + if tokens[0] == "nameserver": self.nameservers.append(tokens[1]) - elif tokens[0] == 'domain': + elif tokens[0] == "domain": self.domain = dns.name.from_text(tokens[1]) # domain and search are exclusive self.search = [] - elif tokens[0] == 'search': + elif tokens[0] == "search": # the last search wins self.search = [] for suffix in tokens[1:]: self.search.append(dns.name.from_text(suffix)) # We don't set domain as it is not used if # len(self.search) > 0 - elif tokens[0] == 'options': + elif tokens[0] == "options": for opt in tokens[1:]: - if opt == 'rotate': + if opt == "rotate": self.rotate = True - elif opt == 'edns0': + elif opt == "edns0": self.use_edns() - elif 'timeout' in opt: + elif "timeout" in opt: try: - self.timeout = int(opt.split(':')[1]) + self.timeout = int(opt.split(":")[1]) except (ValueError, IndexError): pass - elif 'ndots' in opt: + elif "ndots" in opt: try: - self.ndots = int(opt.split(':')[1]) + self.ndots = int(opt.split(":")[1]) except (ValueError, IndexError): pass if len(self.nameservers) == 0: - raise NoResolverConfiguration('no nameservers') + raise NoResolverConfiguration("no nameservers") def read_registry(self) -> None: """Extract resolver configuration from the Windows registry.""" @@ -913,8 +967,12 @@ class BaseResolver: except AttributeError: raise NotImplementedError - def _compute_timeout(self, start: float, lifetime: Optional[float]=None, - errors: Optional[List[ErrorTuple]]=None) -> float: + def _compute_timeout( + self, + start: float, + lifetime: Optional[float] = None, + errors: Optional[List[ErrorTuple]] = None, + ) -> float: lifetime = self.lifetime if lifetime is None else lifetime now = time.time() duration = now - start @@ -933,7 +991,9 @@ class BaseResolver: raise LifetimeTimeout(timeout=duration, errors=errors) return min(lifetime - duration, self.timeout) - def _get_qnames_to_try(self, qname: dns.name.Name, search: Optional[bool]) -> List[dns.name.Name]: + def _get_qnames_to_try( + self, qname: dns.name.Name, search: Optional[bool] + ) -> List[dns.name.Name]: # This is a separate method so we can unit test the search # rules without requiring the Internet. if search is None: @@ -972,8 +1032,12 @@ class BaseResolver: qnames_to_try.append(abs_qname) return qnames_to_try - def use_tsig(self, keyring: Any, keyname: Optional[Union[dns.name.Name, str]]=None, - algorithm: Union[dns.name.Name, str]=dns.tsig.default_algorithm) -> None: + def use_tsig( + self, + keyring: Any, + keyname: Optional[Union[dns.name.Name, str]] = None, + algorithm: Union[dns.name.Name, str] = dns.tsig.default_algorithm, + ) -> None: """Add a TSIG signature to each query. The parameters are passed to ``dns.message.Message.use_tsig()``; @@ -984,9 +1048,13 @@ class BaseResolver: self.keyname = keyname self.keyalgorithm = algorithm - def use_edns(self, edns: Optional[Union[int, bool]]=0, ednsflags: int=0, - payload: int=dns.message.DEFAULT_EDNS_PAYLOAD, - options: Optional[List[dns.edns.Option]]=None) -> None: + def use_edns( + self, + edns: Optional[Union[int, bool]] = 0, + ednsflags: int = 0, + payload: int = dns.message.DEFAULT_EDNS_PAYLOAD, + options: Optional[List[dns.edns.Option]] = None, + ) -> None: """Configure EDNS behavior. *edns*, an ``int``, is the EDNS level to use. Specifying @@ -1037,27 +1105,35 @@ class BaseResolver: for nameserver in nameservers: if not dns.inet.is_address(nameserver): try: - if urlparse(nameserver).scheme != 'https': + if urlparse(nameserver).scheme != "https": raise NotImplementedError except Exception: - raise ValueError(f'nameserver {nameserver} is not an ' - 'IP address or valid https URL') + raise ValueError( + f"nameserver {nameserver} is not an " + "IP address or valid https URL" + ) self._nameservers = nameservers else: - raise ValueError('nameservers must be a list' - ' (not a {})'.format(type(nameservers))) + raise ValueError( + "nameservers must be a list" " (not a {})".format(type(nameservers)) + ) class Resolver(BaseResolver): """DNS stub resolver.""" - def resolve(self, qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, - tcp: bool = False, source: Optional[str] = None, - raise_on_no_answer: bool = True, source_port: int = 0, - lifetime: Optional[float] = None, - search: Optional[bool] = None) -> Answer: # pylint: disable=arguments-differ + def resolve( + self, + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, + search: Optional[bool] = None, + ) -> Answer: # pylint: disable=arguments-differ """Query nameservers to find the answer to the question. The *qname*, *rdtype*, and *rdclass* parameters may be objects @@ -1109,8 +1185,9 @@ class Resolver(BaseResolver): """ - resolution = _Resolution(self, qname, rdtype, rdclass, tcp, - raise_on_no_answer, search) + resolution = _Resolution( + self, qname, rdtype, rdclass, tcp, raise_on_no_answer, search + ) start = time.time() while True: (request, answer) = resolution.next_request() @@ -1127,27 +1204,30 @@ class Resolver(BaseResolver): (nameserver, port, tcp, backoff) = resolution.next_nameserver() if backoff: time.sleep(backoff) - timeout = self._compute_timeout(start, lifetime, - resolution.errors) + timeout = self._compute_timeout(start, lifetime, resolution.errors) try: if dns.inet.is_address(nameserver): if tcp: - response = dns.query.tcp(request, nameserver, - timeout=timeout, - port=port, - source=source, - source_port=source_port) + response = dns.query.tcp( + request, + nameserver, + timeout=timeout, + port=port, + source=source, + source_port=source_port, + ) else: - response = dns.query.udp(request, - nameserver, - timeout=timeout, - port=port, - source=source, - source_port=source_port, - raise_on_truncation=True) + response = dns.query.udp( + request, + nameserver, + timeout=timeout, + port=port, + source=source, + source_port=source_port, + raise_on_truncation=True, + ) else: - response = dns.query.https(request, nameserver, - timeout=timeout) + response = dns.query.https(request, nameserver, timeout=timeout) except Exception as ex: (_, done) = resolution.query_result(None, ex) continue @@ -1159,11 +1239,17 @@ class Resolver(BaseResolver): if answer is not None: return answer - def query(self, qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - tcp: bool=False, source: Optional[str]=None, raise_on_no_answer: bool=True, source_port: int=0, - lifetime: Optional[float]=None) -> Answer: # pragma: no cover + def query( + self, + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, + ) -> Answer: # pragma: no cover """Query nameservers to find the answer to the question. This method calls resolve() with ``search=True``, and is @@ -1171,13 +1257,26 @@ class Resolver(BaseResolver): dnspython. See the documentation for the resolve() method for further details. """ - warnings.warn('please use dns.resolver.Resolver.resolve() instead', - DeprecationWarning, stacklevel=2) - return self.resolve(qname, rdtype, rdclass, tcp, source, - raise_on_no_answer, source_port, lifetime, - True) - - def resolve_address(self, ipaddr: str, *args: Any, **kwargs: Dict[str, Any]) -> Answer: + warnings.warn( + "please use dns.resolver.Resolver.resolve() instead", + DeprecationWarning, + stacklevel=2, + ) + return self.resolve( + qname, + rdtype, + rdclass, + tcp, + source, + raise_on_no_answer, + source_port, + lifetime, + True, + ) + + def resolve_address( + self, ipaddr: str, *args: Any, **kwargs: Dict[str, Any] + ) -> Answer: """Use a resolver to run a reverse query for PTR records. This utilizes the resolve() method to perform a PTR lookup on the @@ -1195,10 +1294,11 @@ class Resolver(BaseResolver): # in the kwargs more than once. modified_kwargs: Dict[str, Any] = {} modified_kwargs.update(kwargs) - modified_kwargs['rdtype'] = dns.rdatatype.PTR - modified_kwargs['rdclass'] = dns.rdataclass.IN - return self.resolve(dns.reversename.from_address(ipaddr), - *args, **modified_kwargs) # type: ignore[arg-type] + modified_kwargs["rdtype"] = dns.rdatatype.PTR + modified_kwargs["rdclass"] = dns.rdataclass.IN + return self.resolve( + dns.reversename.from_address(ipaddr), *args, **modified_kwargs + ) # type: ignore[arg-type] # pylint: disable=redefined-outer-name @@ -1249,11 +1349,17 @@ def reset_default_resolver(): default_resolver = Resolver() -def resolve(qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - tcp: bool=False, source: Optional[str]=None, raise_on_no_answer: bool=True, source_port: int=0, - lifetime: Optional[float]=None, search: Optional[bool]=None) -> Answer: # pragma: no cover +def resolve( + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, + search: Optional[bool] = None, +) -> Answer: # pragma: no cover """Query nameservers to find the answer to the question. @@ -1264,15 +1370,29 @@ def resolve(qname: Union[dns.name.Name, str], parameters. """ - return get_default_resolver().resolve(qname, rdtype, rdclass, tcp, source, - raise_on_no_answer, source_port, - lifetime, search) - -def query(qname: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.A, - rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - tcp: bool=False, source: Optional[str]=None, raise_on_no_answer: bool=True, source_port: int=0, - lifetime: Optional[float]=None) -> Answer: # pragma: no cover + return get_default_resolver().resolve( + qname, + rdtype, + rdclass, + tcp, + source, + raise_on_no_answer, + source_port, + lifetime, + search, + ) + + +def query( + qname: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.A, + rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + tcp: bool = False, + source: Optional[str] = None, + raise_on_no_answer: bool = True, + source_port: int = 0, + lifetime: Optional[float] = None, +) -> Answer: # pragma: no cover """Query nameservers to find the answer to the question. This method calls resolve() with ``search=True``, and is @@ -1280,11 +1400,20 @@ def query(qname: Union[dns.name.Name, str], dnspython. See the documentation for the resolve() method for further details. """ - warnings.warn('please use dns.resolver.resolve() instead', - DeprecationWarning, stacklevel=2) - return resolve(qname, rdtype, rdclass, tcp, source, - raise_on_no_answer, source_port, lifetime, - True) + warnings.warn( + "please use dns.resolver.resolve() instead", DeprecationWarning, stacklevel=2 + ) + return resolve( + qname, + rdtype, + rdclass, + tcp, + source, + raise_on_no_answer, + source_port, + lifetime, + True, + ) def resolve_address(ipaddr: str, *args: Any, **kwargs: Dict[str, Any]) -> Answer: @@ -1307,9 +1436,13 @@ def canonical_name(name: Union[dns.name.Name, str]) -> dns.name.Name: return get_default_resolver().canonical_name(name) -def zone_for_name(name: Union[dns.name.Name, str], rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, - tcp: bool=False, resolver: Optional[Resolver]=None, - lifetime: Optional[float]=None) -> dns.name.Name: +def zone_for_name( + name: Union[dns.name.Name, str], + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + tcp: bool = False, + resolver: Optional[Resolver] = None, + lifetime: Optional[float] = None, +) -> dns.name.Name: """Find the name of the zone which contains the specified name. *name*, an absolute ``dns.name.Name`` or ``str``, the query name. @@ -1356,8 +1489,9 @@ def zone_for_name(name: Union[dns.name.Name, str], rdclass: dns.rdataclass.Rdata rlifetime = 0 else: rlifetime = None - answer = resolver.resolve(name, dns.rdatatype.SOA, rdclass, tcp, - lifetime=rlifetime) + answer = resolver.resolve( + name, dns.rdatatype.SOA, rdclass, tcp, lifetime=rlifetime + ) assert answer.rrset is not None if answer.rrset.name == name: return name @@ -1386,6 +1520,7 @@ def zone_for_name(name: Union[dns.name.Name, str], rdclass: dns.rdataclass.Rdata except dns.name.NoParent: raise NoRootSOA + # # Support for overriding the system resolver for all python code in the # running process. @@ -1405,16 +1540,16 @@ _original_gethostbyname_ex = socket.gethostbyname_ex _original_gethostbyaddr = socket.gethostbyaddr -def _getaddrinfo(host=None, service=None, family=socket.AF_UNSPEC, socktype=0, - proto=0, flags=0): +def _getaddrinfo( + host=None, service=None, family=socket.AF_UNSPEC, socktype=0, proto=0, flags=0 +): if flags & socket.AI_NUMERICHOST != 0: # Short circuit directly into the system's getaddrinfo(). We're # not adding any value in this case, and this avoids infinite loops # because dns.query.* needs to call getaddrinfo() for IPv6 scoping # reasons. We will also do this short circuit below if we # discover that the host is an address literal. - return _original_getaddrinfo(host, service, family, socktype, proto, - flags) + return _original_getaddrinfo(host, service, family, socktype, proto, flags) if flags & (socket.AI_ADDRCONFIG | socket.AI_V4MAPPED) != 0: # Not implemented. We raise a gaierror as opposed to a # NotImplementedError as it helps callers handle errors more @@ -1424,32 +1559,30 @@ def _getaddrinfo(host=None, service=None, family=socket.AF_UNSPEC, socktype=0, # no EAI_SYSTEM on Windows [Issue #416]. We didn't go for # EAI_BADFLAGS as the flags aren't bad, we just don't # implement them. - raise socket.gaierror(socket.EAI_FAIL, - 'Non-recoverable failure in name resolution') + raise socket.gaierror( + socket.EAI_FAIL, "Non-recoverable failure in name resolution" + ) if host is None and service is None: - raise socket.gaierror(socket.EAI_NONAME, 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") v6addrs = [] v4addrs = [] canonical_name = None # pylint: disable=redefined-outer-name # Is host None or an address literal? If so, use the system's # getaddrinfo(). if host is None: - return _original_getaddrinfo(host, service, family, socktype, - proto, flags) + return _original_getaddrinfo(host, service, family, socktype, proto, flags) try: # We don't care about the result of af_for_address(), we're just # calling it so it raises an exception if host is not an IPv4 or # IPv6 address. dns.inet.af_for_address(host) - return _original_getaddrinfo(host, service, family, socktype, - proto, flags) + return _original_getaddrinfo(host, service, family, socktype, proto, flags) except Exception: pass # Something needs resolution! try: if family == socket.AF_INET6 or family == socket.AF_UNSPEC: - v6 = _resolver.resolve(host, dns.rdatatype.AAAA, - raise_on_no_answer=False) + v6 = _resolver.resolve(host, dns.rdatatype.AAAA, raise_on_no_answer=False) # Note that setting host ensures we query the same name # for A as we did for AAAA. (This is just in case search lists # are active by default in the resolver configuration and @@ -1461,20 +1594,18 @@ def _getaddrinfo(host=None, service=None, family=socket.AF_UNSPEC, socktype=0, for rdata in v6.rrset: v6addrs.append(rdata.address) if family == socket.AF_INET or family == socket.AF_UNSPEC: - v4 = _resolver.resolve(host, dns.rdatatype.A, - raise_on_no_answer=False) + v4 = _resolver.resolve(host, dns.rdatatype.A, raise_on_no_answer=False) canonical_name = v4.canonical_name.to_text(True) if v4.rrset is not None: for rdata in v4.rrset: v4addrs.append(rdata.address) except dns.resolver.NXDOMAIN: - raise socket.gaierror(socket.EAI_NONAME, 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") except Exception: # We raise EAI_AGAIN here as the failure may be temporary # (e.g. a timeout) and EAI_SYSTEM isn't defined on Windows. # [Issue #416] - raise socket.gaierror(socket.EAI_AGAIN, - 'Temporary failure in name resolution') + raise socket.gaierror(socket.EAI_AGAIN, "Temporary failure in name resolution") port = None try: # Is it a port literal? @@ -1489,7 +1620,7 @@ def _getaddrinfo(host=None, service=None, family=socket.AF_UNSPEC, socktype=0, except Exception: pass if port is None: - raise socket.gaierror(socket.EAI_NONAME, 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") tuples = [] if socktype == 0: socktypes = [socket.SOCK_DGRAM, socket.SOCK_STREAM] @@ -1498,21 +1629,23 @@ def _getaddrinfo(host=None, service=None, family=socket.AF_UNSPEC, socktype=0, if flags & socket.AI_CANONNAME != 0: cname = canonical_name else: - cname = '' + cname = "" if family == socket.AF_INET6 or family == socket.AF_UNSPEC: for addr in v6addrs: for socktype in socktypes: for proto in _protocols_for_socktype[socktype]: - tuples.append((socket.AF_INET6, socktype, proto, - cname, (addr, port, 0, 0))) + tuples.append( + (socket.AF_INET6, socktype, proto, cname, (addr, port, 0, 0)) + ) if family == socket.AF_INET or family == socket.AF_UNSPEC: for addr in v4addrs: for socktype in socktypes: for proto in _protocols_for_socktype[socktype]: - tuples.append((socket.AF_INET, socktype, proto, - cname, (addr, port))) + tuples.append( + (socket.AF_INET, socktype, proto, cname, (addr, port)) + ) if len(tuples) == 0: - raise socket.gaierror(socket.EAI_NONAME, 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") return tuples @@ -1525,31 +1658,29 @@ def _getnameinfo(sockaddr, flags=0): else: scope = None family = socket.AF_INET - tuples = _getaddrinfo(host, port, family, socket.SOCK_STREAM, - socket.SOL_TCP, 0) + tuples = _getaddrinfo(host, port, family, socket.SOCK_STREAM, socket.SOL_TCP, 0) if len(tuples) > 1: - raise socket.error('sockaddr resolved to multiple addresses') + raise socket.error("sockaddr resolved to multiple addresses") addr = tuples[0][4][0] if flags & socket.NI_DGRAM: - pname = 'udp' + pname = "udp" else: - pname = 'tcp' + pname = "tcp" qname = dns.reversename.from_address(addr) if flags & socket.NI_NUMERICHOST == 0: try: - answer = _resolver.resolve(qname, 'PTR') + answer = _resolver.resolve(qname, "PTR") hostname = answer.rrset[0].target.to_text(True) except (dns.resolver.NXDOMAIN, dns.resolver.NoAnswer): if flags & socket.NI_NAMEREQD: - raise socket.gaierror(socket.EAI_NONAME, - 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") hostname = addr if scope is not None: - hostname += '%' + str(scope) + hostname += "%" + str(scope) else: hostname = addr if scope is not None: - hostname += '%' + str(scope) + hostname += "%" + str(scope) if flags & socket.NI_NUMERICSERV: service = str(port) else: @@ -1576,8 +1707,9 @@ def _gethostbyname(name): def _gethostbyname_ex(name): aliases = [] addresses = [] - tuples = _getaddrinfo(name, 0, socket.AF_INET, socket.SOCK_STREAM, - socket.SOL_TCP, socket.AI_CANONNAME) + tuples = _getaddrinfo( + name, 0, socket.AF_INET, socket.SOCK_STREAM, socket.SOL_TCP, socket.AI_CANONNAME + ) canonical = tuples[0][3] for item in tuples: addresses.append(item[4][0]) @@ -1594,15 +1726,15 @@ def _gethostbyaddr(ip): try: dns.ipv4.inet_aton(ip) except Exception: - raise socket.gaierror(socket.EAI_NONAME, - 'Name or service not known') + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") sockaddr = (ip, 80) family = socket.AF_INET (name, _) = _getnameinfo(sockaddr, socket.NI_NAMEREQD) aliases = [] addresses = [] - tuples = _getaddrinfo(name, 0, family, socket.SOCK_STREAM, socket.SOL_TCP, - socket.AI_CANONNAME) + tuples = _getaddrinfo( + name, 0, family, socket.SOCK_STREAM, socket.SOL_TCP, socket.AI_CANONNAME + ) canonical = tuples[0][3] # We only want to include an address from the tuples if it's the # same as the one we asked about. We do this comparison in binary @@ -1617,7 +1749,7 @@ def _gethostbyaddr(ip): return (canonical, aliases, addresses) -def override_system_resolver(resolver: Optional[Resolver]=None) -> None: +def override_system_resolver(resolver: Optional[Resolver] = None) -> None: """Override the system resolver routines in the socket module with versions which use dnspython's resolver. diff --git a/dns/reversename.py b/dns/reversename.py index c25e77d..eb6a3b6 100644 --- a/dns/reversename.py +++ b/dns/reversename.py @@ -23,12 +23,15 @@ import dns.name import dns.ipv6 import dns.ipv4 -ipv4_reverse_domain = dns.name.from_text('in-addr.arpa.') -ipv6_reverse_domain = dns.name.from_text('ip6.arpa.') +ipv4_reverse_domain = dns.name.from_text("in-addr.arpa.") +ipv6_reverse_domain = dns.name.from_text("ip6.arpa.") -def from_address(text: str, v4_origin: dns.name.Name=ipv4_reverse_domain, - v6_origin: dns.name.Name=ipv6_reverse_domain) -> dns.name.Name: +def from_address( + text: str, + v4_origin: dns.name.Name = ipv4_reverse_domain, + v6_origin: dns.name.Name = ipv6_reverse_domain, +) -> dns.name.Name: """Convert an IPv4 or IPv6 address in textual form into a Name object whose value is the reverse-map domain name of the address. @@ -51,20 +54,22 @@ def from_address(text: str, v4_origin: dns.name.Name=ipv4_reverse_domain, try: v6 = dns.ipv6.inet_aton(text) if dns.ipv6.is_mapped(v6): - parts = ['%d' % byte for byte in v6[12:]] + parts = ["%d" % byte for byte in v6[12:]] origin = v4_origin else: parts = [x for x in str(binascii.hexlify(v6).decode())] origin = v6_origin except Exception: - parts = ['%d' % - byte for byte in dns.ipv4.inet_aton(text)] + parts = ["%d" % byte for byte in dns.ipv4.inet_aton(text)] origin = v4_origin - return dns.name.from_text('.'.join(reversed(parts)), origin=origin) + return dns.name.from_text(".".join(reversed(parts)), origin=origin) -def to_address(name: dns.name.Name, v4_origin: dns.name.Name=ipv4_reverse_domain, - v6_origin: dns.name.Name=ipv6_reverse_domain) -> str: +def to_address( + name: dns.name.Name, + v4_origin: dns.name.Name = ipv4_reverse_domain, + v6_origin: dns.name.Name = ipv6_reverse_domain, +) -> str: """Convert a reverse map domain name into textual address form. *name*, a ``dns.name.Name``, an IPv4 or IPv6 address in reverse-map name @@ -84,7 +89,7 @@ def to_address(name: dns.name.Name, v4_origin: dns.name.Name=ipv4_reverse_domain if name.is_subdomain(v4_origin): name = name.relativize(v4_origin) - text = b'.'.join(reversed(name.labels)) + text = b".".join(reversed(name.labels)) # run through inet_ntoa() to check syntax and make pretty. return dns.ipv4.inet_ntoa(dns.ipv4.inet_aton(text)) elif name.is_subdomain(v6_origin): @@ -92,9 +97,9 @@ def to_address(name: dns.name.Name, v4_origin: dns.name.Name=ipv4_reverse_domain labels = list(reversed(name.labels)) parts = [] for i in range(0, len(labels), 4): - parts.append(b''.join(labels[i:i + 4])) - text = b':'.join(parts) + parts.append(b"".join(labels[i : i + 4])) + text = b":".join(parts) # run through inet_ntoa() to check syntax and make pretty. return dns.ipv6.inet_ntoa(dns.ipv6.inet_aton(text)) else: - raise dns.exception.SyntaxError('unknown reverse-map address family') + raise dns.exception.SyntaxError("unknown reverse-map address family") diff --git a/dns/rrset.py b/dns/rrset.py index bfa4763..4217a04 100644 --- a/dns/rrset.py +++ b/dns/rrset.py @@ -36,12 +36,16 @@ class RRset(dns.rdataset.Rdataset): name. """ - __slots__ = ['name', 'deleting'] - - def __init__(self, name: dns.name.Name, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - deleting: Optional[dns.rdataclass.RdataClass]=None): + __slots__ = ["name", "deleting"] + + def __init__( + self, + name: dns.name.Name, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + deleting: Optional[dns.rdataclass.RdataClass] = None, + ): """Create a new RRset.""" super().__init__(rdclass, rdtype, covers) @@ -56,17 +60,26 @@ class RRset(dns.rdataset.Rdataset): def __repr__(self): if self.covers == 0: - ctext = '' + ctext = "" else: - ctext = '(' + dns.rdatatype.to_text(self.covers) + ')' + ctext = "(" + dns.rdatatype.to_text(self.covers) + ")" if self.deleting is not None: - dtext = ' delete=' + dns.rdataclass.to_text(self.deleting) + dtext = " delete=" + dns.rdataclass.to_text(self.deleting) else: - dtext = '' - return '<DNS ' + str(self.name) + ' ' + \ - dns.rdataclass.to_text(self.rdclass) + ' ' + \ - dns.rdatatype.to_text(self.rdtype) + ctext + dtext + \ - ' RRset: ' + self._rdata_repr() + '>' + dtext = "" + return ( + "<DNS " + + str(self.name) + + " " + + dns.rdataclass.to_text(self.rdclass) + + " " + + dns.rdatatype.to_text(self.rdtype) + + ctext + + dtext + + " RRset: " + + self._rdata_repr() + + ">" + ) def __str__(self): return self.to_text() @@ -79,7 +92,9 @@ class RRset(dns.rdataset.Rdataset): return False return super().__eq__(other) - def match(self, *args: Any, **kwargs: Dict[str, Any]) -> bool: # type: ignore[override] + def match( # type: ignore[override] + self, *args: Any, **kwargs: Dict[str, Any] + ) -> bool: """Does this rrset match the specified attributes? Behaves as :py:func:`full_match()` if the first argument is a @@ -96,9 +111,14 @@ class RRset(dns.rdataset.Rdataset): else: return super().match(*args, **kwargs) # type: ignore[arg-type] - def full_match(self, name: dns.name.Name, rdclass: dns.rdataclass.RdataClass, - rdtype: dns.rdatatype.RdataType, covers: dns.rdatatype.RdataType, - deleting: Optional[dns.rdataclass.RdataClass]=None) -> bool: + def full_match( + self, + name: dns.name.Name, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType, + deleting: Optional[dns.rdataclass.RdataClass] = None, + ) -> bool: """Returns ``True`` if this rrset matches the specified name, class, type, covers, and deletion state. """ @@ -110,7 +130,12 @@ class RRset(dns.rdataset.Rdataset): # pylint: disable=arguments-differ - def to_text(self, origin: Optional[dns.name.Name]=None, relativize: bool=True, **kw) -> str: # type: ignore + def to_text( # type: ignore[override] + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + **kw: Dict[str, Any] + ) -> str: """Convert the RRset into DNS zone file format. See ``dns.name.Name.choose_relativity`` for more information @@ -127,11 +152,17 @@ class RRset(dns.rdataset.Rdataset): to *origin*. """ - return super().to_text(self.name, origin, relativize, - self.deleting, **kw) - - def to_wire(self, file: Any, compress: Optional[dns.name.CompressType]=None, # type: ignore - origin: Optional[dns.name.Name]=None, **kw) -> int: + return super().to_text( + self.name, origin, relativize, self.deleting, **kw # type: ignore + ) + + def to_wire( # type: ignore[override] + self, + file: Any, + compress: Optional[dns.name.CompressType] = None, # type: ignore + origin: Optional[dns.name.Name] = None, + **kw: Dict[str, Any] + ) -> int: """Convert the RRset to wire format. All keyword arguments are passed to ``dns.rdataset.to_wire()``; see @@ -140,8 +171,9 @@ class RRset(dns.rdataset.Rdataset): Returns an ``int``, the number of records emitted. """ - return super().to_wire(self.name, file, compress, origin, - self.deleting, **kw) + return super().to_wire( + self.name, file, compress, origin, self.deleting, **kw # type:ignore + ) # pylint: enable=arguments-differ @@ -153,13 +185,17 @@ class RRset(dns.rdataset.Rdataset): return dns.rdataset.from_rdata_list(self.ttl, list(self)) -def from_text_list(name: Union[dns.name.Name, str], ttl: int, - rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - text_rdatas: Collection[str], - idna_codec: Optional[dns.name.IDNACodec]=None, - origin: Optional[dns.name.Name]=None, relativize: bool=True, - relativize_to: Optional[dns.name.Name]=None) -> RRset: +def from_text_list( + name: Union[dns.name.Name, str], + ttl: int, + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + text_rdatas: Collection[str], + idna_codec: Optional[dns.name.IDNACodec] = None, + origin: Optional[dns.name.Name] = None, + relativize: bool = True, + relativize_to: Optional[dns.name.Name] = None, +) -> RRset: """Create an RRset with the specified name, TTL, class, and type, and with the specified list of rdatas in text format. @@ -185,29 +221,37 @@ def from_text_list(name: Union[dns.name.Name, str], ttl: int, r = RRset(name, the_rdclass, the_rdtype) r.update_ttl(ttl) for t in text_rdatas: - rd = dns.rdata.from_text(r.rdclass, r.rdtype, t, origin, relativize, - relativize_to, idna_codec) + rd = dns.rdata.from_text( + r.rdclass, r.rdtype, t, origin, relativize, relativize_to, idna_codec + ) r.add(rd) return r -def from_text(name: Union[dns.name.Name, str], ttl: int, - rdclass: Union[dns.rdataclass.RdataClass, str], - rdtype: Union[dns.rdatatype.RdataType, str], - *text_rdatas: Any) -> RRset: +def from_text( + name: Union[dns.name.Name, str], + ttl: int, + rdclass: Union[dns.rdataclass.RdataClass, str], + rdtype: Union[dns.rdatatype.RdataType, str], + *text_rdatas: Any +) -> RRset: """Create an RRset with the specified name, TTL, class, and type and with the specified rdatas in text format. Returns a ``dns.rrset.RRset`` object. """ - return from_text_list(name, ttl, rdclass, rdtype, - cast(Collection[str], text_rdatas)) + return from_text_list( + name, ttl, rdclass, rdtype, cast(Collection[str], text_rdatas) + ) -def from_rdata_list(name: Union[dns.name.Name, str], ttl: int, - rdatas: Collection[dns.rdata.Rdata], - idna_codec: Optional[dns.name.IDNACodec]=None) -> RRset: +def from_rdata_list( + name: Union[dns.name.Name, str], + ttl: int, + rdatas: Collection[dns.rdata.Rdata], + idna_codec: Optional[dns.name.IDNACodec] = None, +) -> RRset: """Create an RRset with the specified name and TTL, and with the specified list of rdata objects. @@ -234,7 +278,7 @@ def from_rdata_list(name: Union[dns.name.Name, str], ttl: int, return r -def from_rdata(name: Union[dns.name.Name, str], ttl:int, *rdatas: Any) -> RRset: +def from_rdata(name: Union[dns.name.Name, str], ttl: int, *rdatas: Any) -> RRset: """Create an RRset with the specified name and TTL, and with the specified rdata objects. diff --git a/dns/serial.py b/dns/serial.py index b4d264c..3417299 100644 --- a/dns/serial.py +++ b/dns/serial.py @@ -2,13 +2,14 @@ """Serial Number Arthimetic from RFC 1982""" + class Serial: - def __init__(self, value: int, bits: int=32): - self.value = value % 2 ** bits + def __init__(self, value: int, bits: int = 32): + self.value = value % 2**bits self.bits = bits def __repr__(self): - return f'dns.serial.Serial({self.value}, {self.bits})' + return f"dns.serial.Serial({self.value}, {self.bits})" def __eq__(self, other): if isinstance(other, int): @@ -29,11 +30,11 @@ class Serial: other = Serial(other, self.bits) elif not isinstance(other, Serial) or other.bits != self.bits: return NotImplemented - if self.value < other.value and \ - other.value - self.value < 2 ** (self.bits - 1): + if self.value < other.value and other.value - self.value < 2 ** (self.bits - 1): return True - elif self.value > other.value and \ - self.value - other.value > 2 ** (self.bits - 1): + elif self.value > other.value and self.value - other.value > 2 ** ( + self.bits - 1 + ): return True else: return False @@ -46,11 +47,11 @@ class Serial: other = Serial(other, self.bits) elif not isinstance(other, Serial) or other.bits != self.bits: return NotImplemented - if self.value < other.value and \ - other.value - self.value > 2 ** (self.bits - 1): + if self.value < other.value and other.value - self.value > 2 ** (self.bits - 1): return True - elif self.value > other.value and \ - self.value - other.value < 2 ** (self.bits - 1): + elif self.value > other.value and self.value - other.value < 2 ** ( + self.bits - 1 + ): return True else: return False @@ -69,7 +70,7 @@ class Serial: if abs(delta) > (2 ** (self.bits - 1) - 1): raise ValueError v += delta - v = v % 2 ** self.bits + v = v % 2**self.bits return Serial(v, self.bits) def __iadd__(self, other): @@ -83,7 +84,7 @@ class Serial: if abs(delta) > (2 ** (self.bits - 1) - 1): raise ValueError v += delta - v = v % 2 ** self.bits + v = v % 2**self.bits self.value = v return self @@ -98,7 +99,7 @@ class Serial: if abs(delta) > (2 ** (self.bits - 1) - 1): raise ValueError v -= delta - v = v % 2 ** self.bits + v = v % 2**self.bits return Serial(v, self.bits) def __isub__(self, other): @@ -112,6 +113,6 @@ class Serial: if abs(delta) > (2 ** (self.bits - 1) - 1): raise ValueError v -= delta - v = v % 2 ** self.bits + v = v % 2**self.bits self.value = v return self @@ -28,7 +28,7 @@ class Set: ability is widely used in dnspython applications. """ - __slots__ = ['items'] + __slots__ = ["items"] def __init__(self, items=None): """Initialize the set. @@ -47,15 +47,13 @@ class Set: return "dns.set.Set(%s)" % repr(list(self.items.keys())) def add(self, item): - """Add an item to the set. - """ + """Add an item to the set.""" if item not in self.items: self.items[item] = None def remove(self, item): - """Remove an item from the set. - """ + """Remove an item from the set.""" try: del self.items[item] @@ -63,8 +61,7 @@ class Set: raise ValueError def discard(self, item): - """Remove an item from the set if present. - """ + """Remove an item from the set if present.""" self.items.pop(item, None) @@ -73,7 +70,7 @@ class Set: (k, _) = self.items.popitem() return k - def _clone(self) -> 'Set': + def _clone(self) -> "Set": """Make a (shallow) copy of the set. There is a 'clone protocol' that subclasses of this class @@ -86,8 +83,8 @@ class Set: subclasses. """ - if hasattr(self, '_clone_class'): - cls = self._clone_class # type: ignore + if hasattr(self, "_clone_class"): + cls = self._clone_class # type: ignore else: cls = self.__class__ obj = cls.__new__(cls) @@ -96,14 +93,12 @@ class Set: return obj def __copy__(self): - """Make a (shallow) copy of the set. - """ + """Make a (shallow) copy of the set.""" return self._clone() def copy(self): - """Make a (shallow) copy of the set. - """ + """Make a (shallow) copy of the set.""" return self._clone() @@ -113,7 +108,7 @@ class Set: """ if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") if self is other: # lgtm[py/comparison-using-is] return for item in other.items: @@ -125,7 +120,7 @@ class Set: """ if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") if self is other: # lgtm[py/comparison-using-is] return # we make a copy of the list so that we can remove items from @@ -140,7 +135,7 @@ class Set: """ if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") if self is other: # lgtm[py/comparison-using-is] self.items.clear() else: @@ -151,7 +146,7 @@ class Set: """Update the set, retaining only elements unique to both sets.""" if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") if self is other: # lgtm[py/comparison-using-is] self.items.clear() else: @@ -285,7 +280,7 @@ class Set: """ if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") for item in self.items: if item not in other.items: return False @@ -298,7 +293,7 @@ class Set: """ if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") for item in other.items: if item not in self.items: return False @@ -306,7 +301,7 @@ class Set: def isdisjoint(self, other): if not isinstance(other, Set): - raise ValueError('other must be a Set instance') + raise ValueError("other must be a Set instance") for item in other.items: if item in self.items: return False diff --git a/dns/tokenizer.py b/dns/tokenizer.py index 275861c..0551578 100644 --- a/dns/tokenizer.py +++ b/dns/tokenizer.py @@ -17,7 +17,7 @@ """Tokenize DNS zone file format""" -from typing import Any, Optional, List, Tuple, Union +from typing import Any, Optional, List, Tuple import io import sys @@ -26,7 +26,7 @@ import dns.exception import dns.name import dns.ttl -_DELIMITERS = {' ', '\t', '\n', ';', '(', ')', '"'} +_DELIMITERS = {" ", "\t", "\n", ";", "(", ")", '"'} _QUOTING_DELIMITERS = {'"'} EOF = 0 @@ -50,8 +50,13 @@ class Token: has_escape: Does the token value contain escapes? """ - def __init__(self, ttype: int, value: Any='', has_escape: bool=False, - comment: Optional[str]=None): + def __init__( + self, + ttype: int, + value: Any = "", + has_escape: bool = False, + comment: Optional[str] = None, + ): """Initialize a token instance.""" self.ttype = ttype @@ -86,28 +91,26 @@ class Token: def __eq__(self, other): if not isinstance(other, Token): return False - return (self.ttype == other.ttype and - self.value == other.value) + return self.ttype == other.ttype and self.value == other.value def __ne__(self, other): if not isinstance(other, Token): return True - return (self.ttype != other.ttype or - self.value != other.value) + return self.ttype != other.ttype or self.value != other.value def __str__(self): return '%d "%s"' % (self.ttype, self.value) - def unescape(self) -> 'Token': + def unescape(self) -> "Token": if not self.has_escape: return self - unescaped = '' + unescaped = "" l = len(self.value) i = 0 while i < l: c = self.value[i] i += 1 - if c == '\\': + if c == "\\": if i >= l: # pragma: no cover (can't happen via get()) raise dns.exception.UnexpectedEnd c = self.value[i] @@ -130,7 +133,7 @@ class Token: unescaped += c return Token(self.ttype, unescaped) - def unescape_to_bytes(self) -> 'Token': + def unescape_to_bytes(self) -> "Token": # We used to use unescape() for TXT-like records, but this # caused problems as we'd process DNS escapes into Unicode code # points instead of byte values, and then a to_text() of the @@ -155,13 +158,13 @@ class Token: # # foo\226\128\139bar # - unescaped = b'' + unescaped = b"" l = len(self.value) i = 0 while i < l: c = self.value[i] i += 1 - if c == '\\': + if c == "\\": if i >= l: # pragma: no cover (can't happen via get()) raise dns.exception.UnexpectedEnd c = self.value[i] @@ -180,7 +183,7 @@ class Token: codepoint = int(c) * 100 + int(c2) * 10 + int(c3) if codepoint > 255: raise dns.exception.SyntaxError - unescaped += b'%c' % (codepoint) + unescaped += b"%c" % (codepoint) else: # Note that as mentioned above, if c is a Unicode # code point outside of the ASCII range, then this @@ -226,8 +229,12 @@ class Tokenizer: encoder/decoder is used. """ - def __init__(self, f: Any=sys.stdin, filename: Optional[str]=None, - idna_codec: Optional[dns.name.IDNACodec]=None): + def __init__( + self, + f: Any = sys.stdin, + filename: Optional[str] = None, + idna_codec: Optional[dns.name.IDNACodec] = None, + ): """Initialize a tokenizer instance. f: The file to tokenize. The default is sys.stdin. @@ -245,17 +252,17 @@ class Tokenizer: if isinstance(f, str): f = io.StringIO(f) if filename is None: - filename = '<string>' + filename = "<string>" elif isinstance(f, bytes): f = io.StringIO(f.decode()) if filename is None: - filename = '<string>' + filename = "<string>" else: if filename is None: if f is sys.stdin: - filename = '<stdin>' + filename = "<stdin>" else: - filename = '<file>' + filename = "<file>" self.file = f self.ungotten_char: Optional[str] = None self.ungotten_token: Optional[Token] = None @@ -272,17 +279,16 @@ class Tokenizer: self.idna_codec = idna_codec def _get_char(self) -> str: - """Read a character from input. - """ + """Read a character from input.""" if self.ungotten_char is None: if self.eof: - c = '' + c = "" else: c = self.file.read(1) - if c == '': + if c == "": self.eof = True - elif c == '\n': + elif c == "\n": self.line_number += 1 else: c = self.ungotten_char @@ -328,13 +334,13 @@ class Tokenizer: skipped = 0 while True: c = self._get_char() - if c != ' ' and c != '\t': - if (c != '\n') or not self.multiline: + if c != " " and c != "\t": + if (c != "\n") or not self.multiline: self._unget_char(c) return skipped skipped += 1 - def get(self, want_leading: bool=False, want_comment: bool=False) -> Token: + def get(self, want_leading: bool = False, want_comment: bool = False) -> Token: """Get the next token. want_leading: If True, return a WHITESPACE token if the @@ -363,21 +369,21 @@ class Tokenizer: return utoken skipped = self.skip_whitespace() if want_leading and skipped > 0: - return Token(WHITESPACE, ' ') - token = '' + return Token(WHITESPACE, " ") + token = "" ttype = IDENTIFIER has_escape = False while True: c = self._get_char() - if c == '' or c in self.delimiters: - if c == '' and self.quoting: + if c == "" or c in self.delimiters: + if c == "" and self.quoting: raise dns.exception.UnexpectedEnd - if token == '' and ttype != QUOTED_STRING: - if c == '(': + if token == "" and ttype != QUOTED_STRING: + if c == "(": self.multiline += 1 self.skip_whitespace() continue - elif c == ')': + elif c == ")": if self.multiline <= 0: raise dns.exception.SyntaxError self.multiline -= 1 @@ -394,28 +400,29 @@ class Tokenizer: self.delimiters = _DELIMITERS self.skip_whitespace() continue - elif c == '\n': - return Token(EOL, '\n') - elif c == ';': + elif c == "\n": + return Token(EOL, "\n") + elif c == ";": while 1: c = self._get_char() - if c == '\n' or c == '': + if c == "\n" or c == "": break token += c if want_comment: self._unget_char(c) return Token(COMMENT, token) - elif c == '': + elif c == "": if self.multiline: raise dns.exception.SyntaxError( - 'unbalanced parentheses') + "unbalanced parentheses" + ) return Token(EOF, comment=token) elif self.multiline: self.skip_whitespace() - token = '' + token = "" continue else: - return Token(EOL, '\n', comment=token) + return Token(EOL, "\n", comment=token) else: # This code exists in case we ever want a # delimiter to be returned. It never produces @@ -425,9 +432,9 @@ class Tokenizer: else: self._unget_char(c) break - elif self.quoting and c == '\n': - raise dns.exception.SyntaxError('newline in quoted string') - elif c == '\\': + elif self.quoting and c == "\n": + raise dns.exception.SyntaxError("newline in quoted string") + elif c == "\\": # # It's an escape. Put it and the next character into # the token; it will be checked later for goodness. @@ -435,12 +442,12 @@ class Tokenizer: token += c has_escape = True c = self._get_char() - if c == '' or (c == '\n' and not self.quoting): + if c == "" or (c == "\n" and not self.quoting): raise dns.exception.UnexpectedEnd token += c - if token == '' and ttype != QUOTED_STRING: + if token == "" and ttype != QUOTED_STRING: if self.multiline: - raise dns.exception.SyntaxError('unbalanced parentheses') + raise dns.exception.SyntaxError("unbalanced parentheses") ttype = EOF return Token(ttype, token, has_escape) @@ -478,7 +485,7 @@ class Tokenizer: # Helpers - def get_int(self, base: int=10) -> int: + def get_int(self, base: int = 10) -> int: """Read the next token and interpret it as an unsigned integer. Raises dns.exception.SyntaxError if not an unsigned integer. @@ -488,9 +495,9 @@ class Tokenizer: token = self.get().unescape() if not token.is_identifier(): - raise dns.exception.SyntaxError('expecting an identifier') + raise dns.exception.SyntaxError("expecting an identifier") if not token.value.isdigit(): - raise dns.exception.SyntaxError('expecting an integer') + raise dns.exception.SyntaxError("expecting an integer") return int(token.value, base) def get_uint8(self) -> int: @@ -505,10 +512,11 @@ class Tokenizer: value = self.get_int() if value < 0 or value > 255: raise dns.exception.SyntaxError( - '%d is not an unsigned 8-bit integer' % value) + "%d is not an unsigned 8-bit integer" % value + ) return value - def get_uint16(self, base: int=10) -> int: + def get_uint16(self, base: int = 10) -> int: """Read the next token and interpret it as a 16-bit unsigned integer. @@ -521,13 +529,15 @@ class Tokenizer: if value < 0 or value > 65535: if base == 8: raise dns.exception.SyntaxError( - '%o is not an octal unsigned 16-bit integer' % value) + "%o is not an octal unsigned 16-bit integer" % value + ) else: raise dns.exception.SyntaxError( - '%d is not an unsigned 16-bit integer' % value) + "%d is not an unsigned 16-bit integer" % value + ) return value - def get_uint32(self, base: int=10) -> int: + def get_uint32(self, base: int = 10) -> int: """Read the next token and interpret it as a 32-bit unsigned integer. @@ -539,10 +549,11 @@ class Tokenizer: value = self.get_int(base=base) if value < 0 or value > 4294967295: raise dns.exception.SyntaxError( - '%d is not an unsigned 32-bit integer' % value) + "%d is not an unsigned 32-bit integer" % value + ) return value - def get_uint48(self, base: int=10) -> int: + def get_uint48(self, base: int = 10) -> int: """Read the next token and interpret it as a 48-bit unsigned integer. @@ -554,10 +565,11 @@ class Tokenizer: value = self.get_int(base=base) if value < 0 or value > 281474976710655: raise dns.exception.SyntaxError( - '%d is not an unsigned 48-bit integer' % value) + "%d is not an unsigned 48-bit integer" % value + ) return value - def get_string(self, max_length: Optional[int]=None) -> str: + def get_string(self, max_length: Optional[int] = None) -> str: """Read the next token and interpret it as a string. Raises dns.exception.SyntaxError if not a string. @@ -569,7 +581,7 @@ class Tokenizer: token = self.get().unescape() if not (token.is_identifier() or token.is_quoted_string()): - raise dns.exception.SyntaxError('expecting a string') + raise dns.exception.SyntaxError("expecting a string") if max_length and len(token.value) > max_length: raise dns.exception.SyntaxError("string too long") return token.value @@ -584,10 +596,10 @@ class Tokenizer: token = self.get().unescape() if not token.is_identifier(): - raise dns.exception.SyntaxError('expecting an identifier') + raise dns.exception.SyntaxError("expecting an identifier") return token.value - def get_remaining(self, max_tokens: Optional[int]=None) -> List[Token]: + def get_remaining(self, max_tokens: Optional[int] = None) -> List[Token]: """Return the remaining tokens on the line, until an EOL or EOF is seen. max_tokens: If not None, stop after this number of tokens. @@ -606,7 +618,7 @@ class Tokenizer: break return tokens - def concatenate_remaining_identifiers(self, allow_empty: bool=False) -> str: + def concatenate_remaining_identifiers(self, allow_empty: bool = False) -> str: """Read the remaining tokens on the line, which should be identifiers. Raises dns.exception.SyntaxError if there are no remaining tokens, @@ -628,11 +640,16 @@ class Tokenizer: raise dns.exception.SyntaxError s += token.value if not (allow_empty or s): - raise dns.exception.SyntaxError('expecting another identifier') + raise dns.exception.SyntaxError("expecting another identifier") return s - def as_name(self, token: Token, origin: Optional[dns.name.Name]=None, - relativize: bool=False, relativize_to: Optional[dns.name.Name]=None) -> dns.name.Name: + def as_name( + self, + token: Token, + origin: Optional[dns.name.Name] = None, + relativize: bool = False, + relativize_to: Optional[dns.name.Name] = None, + ) -> dns.name.Name: """Try to interpret the token as a DNS name. Raises dns.exception.SyntaxError if not a name. @@ -640,12 +657,16 @@ class Tokenizer: Returns a dns.name.Name. """ if not token.is_identifier(): - raise dns.exception.SyntaxError('expecting an identifier') + raise dns.exception.SyntaxError("expecting an identifier") name = dns.name.from_text(token.value, origin, self.idna_codec) return name.choose_relativity(relativize_to or origin, relativize) - def get_name(self, origin: Optional[dns.name.Name]=None, relativize: bool=False, - relativize_to: Optional[dns.name.Name]=None) -> dns.name.Name: + def get_name( + self, + origin: Optional[dns.name.Name] = None, + relativize: bool = False, + relativize_to: Optional[dns.name.Name] = None, + ) -> dns.name.Name: """Read the next token and interpret it as a DNS name. Raises dns.exception.SyntaxError if not a name. @@ -666,8 +687,8 @@ class Tokenizer: token = self.get() if not token.is_eol_or_eof(): raise dns.exception.SyntaxError( - 'expected EOL or EOF, got %d "%s"' % (token.ttype, - token.value)) + 'expected EOL or EOF, got %d "%s"' % (token.ttype, token.value) + ) return token def get_eol(self) -> str: @@ -684,5 +705,5 @@ class Tokenizer: token = self.get().unescape() if not token.is_identifier(): - raise dns.exception.SyntaxError('expecting an identifier') + raise dns.exception.SyntaxError("expecting an identifier") return dns.ttl.from_text(token.value) diff --git a/dns/transaction.py b/dns/transaction.py index 1e97c75..b0429df 100644 --- a/dns/transaction.py +++ b/dns/transaction.py @@ -16,11 +16,11 @@ import dns.ttl class TransactionManager: - def reader(self) -> 'Transaction': + def reader(self) -> "Transaction": """Begin a read-only transaction.""" raise NotImplementedError # pragma: no cover - def writer(self, replacement: bool=False) -> 'Transaction': + def writer(self, replacement: bool = False) -> "Transaction": """Begin a writable transaction. *replacement*, a ``bool``. If `True`, the content of the @@ -30,7 +30,9 @@ class TransactionManager: """ raise NotImplementedError # pragma: no cover - def origin_information(self) -> Tuple[Optional[dns.name.Name], bool, Optional[dns.name.Name]]: + def origin_information( + self, + ) -> Tuple[Optional[dns.name.Name], bool, Optional[dns.name.Name]]: """Returns a tuple (absolute_origin, relativize, effective_origin) @@ -56,13 +58,11 @@ class TransactionManager: raise NotImplementedError # pragma: no cover def get_class(self) -> dns.rdataclass.RdataClass: - """The class of the transaction manager. - """ + """The class of the transaction manager.""" raise NotImplementedError # pragma: no cover def from_wire_origin(self) -> Optional[dns.name.Name]: - """Origin to use in from_wire() calls. - """ + """Origin to use in from_wire() calls.""" (absolute_origin, relativize, _) = self.origin_information() if relativize: return absolute_origin @@ -87,39 +87,51 @@ def _ensure_immutable_rdataset(rdataset): return rdataset return dns.rdataset.ImmutableRdataset(rdataset) + def _ensure_immutable_node(node): if node is None or node.is_immutable(): return node return dns.node.ImmutableNode(node) -CheckPutRdatasetType = Callable[['Transaction', dns.name.Name, dns.rdataset.Rdataset], None] -CheckDeleteRdatasetType = Callable[['Transaction', dns.name.Name, - dns.rdatatype.RdataType, dns.rdatatype.RdataType], None] -CheckDeleteNameType = Callable[['Transaction', dns.name.Name], None] +CheckPutRdatasetType = Callable[ + ["Transaction", dns.name.Name, dns.rdataset.Rdataset], None +] +CheckDeleteRdatasetType = Callable[ + ["Transaction", dns.name.Name, dns.rdatatype.RdataType, dns.rdatatype.RdataType], + None, +] +CheckDeleteNameType = Callable[["Transaction", dns.name.Name], None] class Transaction: - - def __init__(self, manager: TransactionManager, replacement: bool=False, read_only: bool=False): + def __init__( + self, + manager: TransactionManager, + replacement: bool = False, + read_only: bool = False, + ): self.manager = manager self.replacement = replacement self.read_only = read_only self._ended = False - self._check_put_rdataset: List[CheckPutRdatasetType]= [] + self._check_put_rdataset: List[CheckPutRdatasetType] = [] self._check_delete_rdataset: List[CheckDeleteRdatasetType] = [] self._check_delete_name: List[CheckDeleteNameType] = [] # # This is the high level API # - # Note that we currently use non-immutable types in the return type signature to avoid - # covariance problems, e.g. if the caller has a List[Rdataset], mypy will be unhappy if we - # return an ImmutableRdataset. - - def get(self, name: Optional[Union[dns.name.Name,str]], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) -> dns.rdataset.Rdataset: + # Note that we currently use non-immutable types in the return type signature to + # avoid covariance problems, e.g. if the caller has a List[Rdataset], mypy will be + # unhappy if we return an ImmutableRdataset. + + def get( + self, + name: Optional[Union[dns.name.Name, str]], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> dns.rdataset.Rdataset: """Return the rdataset associated with *name*, *rdtype*, and *covers*, or `None` if not found. @@ -232,7 +244,12 @@ class Transaction: name = dns.name.from_text(name, None) return self._name_exists(name) - def update_serial(self, value: int=1, relative: bool=True, name: dns.name.Name=dns.name.empty) -> None: + def update_serial( + self, + value: int = 1, + relative: bool = True, + name: dns.name.Name = dns.name.empty, + ) -> None: """Update the serial number. *value*, an `int`, is an increment if *relative* is `True`, or the @@ -246,11 +263,10 @@ class Transaction: """ self._check_ended() if value < 0: - raise ValueError('negative update_serial() value') + raise ValueError("negative update_serial() value") if isinstance(name, str): name = dns.name.from_text(name, None) - rdataset = self._get_rdataset(name, dns.rdatatype.SOA, - dns.rdatatype.NONE) + rdataset = self._get_rdataset(name, dns.rdatatype.SOA, dns.rdatatype.NONE) if rdataset is None or len(rdataset) == 0: raise KeyError if relative: @@ -347,7 +363,7 @@ class Transaction: def _raise_if_not_empty(self, method, args): if len(args) != 0: - raise TypeError(f'extra parameters to {method}') + raise TypeError(f"extra parameters to {method}") def _rdataset_from_args(self, method, deleting, args): try: @@ -363,29 +379,29 @@ class Transaction: if isinstance(arg, int): ttl = arg if ttl > dns.ttl.MAX_TTL: - raise ValueError(f'{method}: TTL value too big') + raise ValueError(f"{method}: TTL value too big") else: - raise TypeError(f'{method}: expected a TTL') + raise TypeError(f"{method}: expected a TTL") arg = args.popleft() if isinstance(arg, dns.rdata.Rdata): rdataset = dns.rdataset.from_rdata(ttl, arg) else: - raise TypeError(f'{method}: expected an Rdata') + raise TypeError(f"{method}: expected an Rdata") return rdataset except IndexError: if deleting: return None else: # reraise - raise TypeError(f'{method}: expected more arguments') + raise TypeError(f"{method}: expected more arguments") def _add(self, replace, args): try: args = collections.deque(args) if replace: - method = 'replace()' + method = "replace()" else: - method = 'add()' + method = "add()" arg = args.popleft() if isinstance(arg, str): arg = dns.name.from_text(arg, None) @@ -399,44 +415,45 @@ class Transaction: # same and can't be stored in nodes, so convert. rdataset = rrset.to_rdataset() else: - raise TypeError(f'{method} requires a name or RRset ' + - 'as the first argument') + raise TypeError( + f"{method} requires a name or RRset " + "as the first argument" + ) if rdataset.rdclass != self.manager.get_class(): - raise ValueError(f'{method} has objects of wrong RdataClass') + raise ValueError(f"{method} has objects of wrong RdataClass") if rdataset.rdtype == dns.rdatatype.SOA: (_, _, origin) = self._origin_information() if name != origin: - raise ValueError(f'{method} has non-origin SOA') + raise ValueError(f"{method} has non-origin SOA") self._raise_if_not_empty(method, args) if not replace: - existing = self._get_rdataset(name, rdataset.rdtype, - rdataset.covers) + existing = self._get_rdataset(name, rdataset.rdtype, rdataset.covers) if existing is not None: if isinstance(existing, dns.rdataset.ImmutableRdataset): - trds = dns.rdataset.Rdataset(existing.rdclass, - existing.rdtype, - existing.covers) + trds = dns.rdataset.Rdataset( + existing.rdclass, existing.rdtype, existing.covers + ) trds.update(existing) existing = trds rdataset = existing.union(rdataset) self._checked_put_rdataset(name, rdataset) except IndexError: - raise TypeError(f'not enough parameters to {method}') + raise TypeError(f"not enough parameters to {method}") def _delete(self, exact, args): try: args = collections.deque(args) if exact: - method = 'delete_exact()' + method = "delete_exact()" else: - method = 'delete()' + method = "delete()" arg = args.popleft() if isinstance(arg, str): arg = dns.name.from_text(arg, None) if isinstance(arg, dns.name.Name): name = arg - if len(args) > 0 and (isinstance(args[0], int) or - isinstance(args[0], str)): + if len(args) > 0 and ( + isinstance(args[0], int) or isinstance(args[0], str) + ): # deleting by type and (optionally) covers rdtype = dns.rdatatype.RdataType.make(args.popleft()) if len(args) > 0: @@ -447,7 +464,7 @@ class Transaction: existing = self._get_rdataset(name, rdtype, covers) if existing is None: if exact: - raise DeleteNotExact(f'{method}: missing rdataset') + raise DeleteNotExact(f"{method}: missing rdataset") else: self._delete_rdataset(name, rdtype, covers) return @@ -457,34 +474,34 @@ class Transaction: rdataset = arg # rrsets are also rdatasets name = rdataset.name else: - raise TypeError(f'{method} requires a name or RRset ' + - 'as the first argument') + raise TypeError( + f"{method} requires a name or RRset " + "as the first argument" + ) self._raise_if_not_empty(method, args) if rdataset: if rdataset.rdclass != self.manager.get_class(): - raise ValueError(f'{method} has objects of wrong ' - 'RdataClass') - existing = self._get_rdataset(name, rdataset.rdtype, - rdataset.covers) + raise ValueError(f"{method} has objects of wrong " "RdataClass") + existing = self._get_rdataset(name, rdataset.rdtype, rdataset.covers) if existing is not None: if exact: intersection = existing.intersection(rdataset) if intersection != rdataset: - raise DeleteNotExact(f'{method}: missing rdatas') + raise DeleteNotExact(f"{method}: missing rdatas") rdataset = existing.difference(rdataset) if len(rdataset) == 0: - self._checked_delete_rdataset(name, rdataset.rdtype, - rdataset.covers) + self._checked_delete_rdataset( + name, rdataset.rdtype, rdataset.covers + ) else: self._checked_put_rdataset(name, rdataset) elif exact: - raise DeleteNotExact(f'{method}: missing rdataset') + raise DeleteNotExact(f"{method}: missing rdataset") else: if exact and not self._name_exists(name): - raise DeleteNotExact(f'{method}: name not known') + raise DeleteNotExact(f"{method}: name not known") self._checked_delete_name(name) except IndexError: - raise TypeError(f'not enough parameters to {method}') + raise TypeError(f"not enough parameters to {method}") def _check_ended(self): if self._ended: @@ -590,8 +607,7 @@ class Transaction: raise NotImplementedError # pragma: no cover def _iterate_rdatasets(self): - """Return an iterator that yields (name, rdataset) tuples. - """ + """Return an iterator that yields (name, rdataset) tuples.""" raise NotImplementedError # pragma: no cover def _get_node(self, name): diff --git a/dns/tsig.py b/dns/tsig.py index 50b2d47..b3f5251 100644 --- a/dns/tsig.py +++ b/dns/tsig.py @@ -27,6 +27,7 @@ import dns.rdataclass import dns.name import dns.rcode + class BadTime(dns.exception.DNSException): """The current time is not within the TSIG's validity time.""" @@ -97,10 +98,11 @@ class GSSTSig: In order to avoid a direct GSSAPI dependency, the keyring holds a ref to the GSSAPI object required, rather than the key itself. """ + def __init__(self, gssapi_context): self.gssapi_context = gssapi_context - self.data = b'' - self.name = 'gss-tsig' + self.data = b"" + self.name = "gss-tsig" def update(self, data): self.data += data @@ -139,9 +141,9 @@ class GSSTSigAdapter: # client to complete the GSSAPI negotiation before attempting # to verify the signed response to a TKEY message exchange try: - rrset = message.find_rrset(message.answer, keyname, - dns.rdataclass.ANY, - dns.rdatatype.TKEY) + rrset = message.find_rrset( + message.answer, keyname, dns.rdataclass.ANY, dns.rdatatype.TKEY + ) if rrset: token = rrset[0].key gssapi_context = key.secret @@ -172,8 +174,9 @@ class HMACTSig: try: hashinfo = self._hashes[algorithm] except KeyError: - raise NotImplementedError(f"TSIG algorithm {algorithm} " + - "is not supported") + raise NotImplementedError( + f"TSIG algorithm {algorithm} " + "is not supported" + ) # create the HMAC context if isinstance(hashinfo, tuple): @@ -184,7 +187,7 @@ class HMACTSig: self.size = None self.name = self.hmac_context.name if self.size: - self.name += f'-{self.size}' + self.name += f"-{self.size}" def update(self, data): return self.hmac_context.update(data) @@ -203,8 +206,7 @@ class HMACTSig: raise BadSignature -def _digest(wire, key, rdata, time=None, request_mac=None, ctx=None, - multi=None): +def _digest(wire, key, rdata, time=None, request_mac=None, ctx=None, multi=None): """Return a context containing the TSIG rdata for the input parameters @rtype: dns.tsig.HMACTSig or dns.tsig.GSSTSig object @raises ValueError: I{other_data} is too long @@ -215,25 +217,25 @@ def _digest(wire, key, rdata, time=None, request_mac=None, ctx=None, if first: ctx = get_context(key) if request_mac: - ctx.update(struct.pack('!H', len(request_mac))) + ctx.update(struct.pack("!H", len(request_mac))) ctx.update(request_mac) - ctx.update(struct.pack('!H', rdata.original_id)) + ctx.update(struct.pack("!H", rdata.original_id)) ctx.update(wire[2:]) if first: ctx.update(key.name.to_digestable()) - ctx.update(struct.pack('!H', dns.rdataclass.ANY)) - ctx.update(struct.pack('!I', 0)) + ctx.update(struct.pack("!H", dns.rdataclass.ANY)) + ctx.update(struct.pack("!I", 0)) if time is None: time = rdata.time_signed - upper_time = (time >> 32) & 0xffff - lower_time = time & 0xffffffff - time_encoded = struct.pack('!HIH', upper_time, lower_time, rdata.fudge) + upper_time = (time >> 32) & 0xFFFF + lower_time = time & 0xFFFFFFFF + time_encoded = struct.pack("!HIH", upper_time, lower_time, rdata.fudge) other_len = len(rdata.other) if other_len > 65535: - raise ValueError('TSIG Other Data is > 65535 bytes') + raise ValueError("TSIG Other Data is > 65535 bytes") if first: ctx.update(key.algorithm.to_digestable() + time_encoded) - ctx.update(struct.pack('!HH', rdata.error, other_len) + rdata.other) + ctx.update(struct.pack("!HH", rdata.error, other_len) + rdata.other) else: ctx.update(time_encoded) return ctx @@ -246,7 +248,7 @@ def _maybe_start_digest(key, mac, multi): """ if multi: ctx = get_context(key) - ctx.update(struct.pack('!H', len(mac))) + ctx.update(struct.pack("!H", len(mac))) ctx.update(mac) return ctx else: @@ -269,8 +271,9 @@ def sign(wire, key, rdata, time=None, request_mac=None, ctx=None, multi=False): return (tsig, _maybe_start_digest(key, mac, multi)) -def validate(wire, key, owner, rdata, now, request_mac, tsig_start, ctx=None, - multi=False): +def validate( + wire, key, owner, rdata, now, request_mac, tsig_start, ctx=None, multi=False +): """Validate the specified TSIG rdata against the other input parameters. @raises FormError: The TSIG is badly formed. @@ -294,7 +297,7 @@ def validate(wire, key, owner, rdata, now, request_mac, tsig_start, ctx=None, elif rdata.error == dns.rcode.BADTRUNC: raise PeerBadTruncation else: - raise PeerError('unknown TSIG error code %d' % rdata.error) + raise PeerError("unknown TSIG error code %d" % rdata.error) if abs(rdata.time_signed - now) > rdata.fudge: raise BadTime if key.name != owner: @@ -332,14 +335,15 @@ class Key: self.algorithm = algorithm def __eq__(self, other): - return (isinstance(other, Key) and - self.name == other.name and - self.secret == other.secret and - self.algorithm == other.algorithm) + return ( + isinstance(other, Key) + and self.name == other.name + and self.secret == other.secret + and self.algorithm == other.algorithm + ) def __repr__(self): - r = f"<DNS key name='{self.name}', " + \ - f"algorithm='{self.algorithm}'" + r = f"<DNS key name='{self.name}', " + f"algorithm='{self.algorithm}'" if self.algorithm != GSS_TSIG: r += f", secret='{base64.b64encode(self.secret).decode()}'" r += ">" diff --git a/dns/tsigkeyring.py b/dns/tsigkeyring.py index ed11794..6adba28 100644 --- a/dns/tsigkeyring.py +++ b/dns/tsigkeyring.py @@ -51,8 +51,10 @@ def to_text(keyring: Dict[dns.name.Name, Any]) -> Dict[str, Any]: @rtype: dict""" textring = {} + def b64encode(secret): return base64.encodebytes(secret).decode().rstrip() + for (name, key) in keyring.items(): tname = name.to_text() if isinstance(key, bytes): @@ -62,15 +62,15 @@ def from_text(text: str) -> int: if need_digit: raise BadTTL c = c.lower() - if c == 'w': + if c == "w": total += current * 604800 - elif c == 'd': + elif c == "d": total += current * 86400 - elif c == 'h': + elif c == "h": total += current * 3600 - elif c == 'm': + elif c == "m": total += current * 60 - elif c == 's': + elif c == "s": total += current else: raise BadTTL("unknown unit '%s'" % c) @@ -89,4 +89,4 @@ def make(value: Union[int, str]) -> int: elif isinstance(value, str): return dns.ttl.from_text(value) else: - raise ValueError('cannot convert value to TTL') + raise ValueError("cannot convert value to TTL") diff --git a/dns/update.py b/dns/update.py index eb7b936..91c8aa4 100644 --- a/dns/update.py +++ b/dns/update.py @@ -31,6 +31,7 @@ import dns.tsig class UpdateSection(dns.enum.IntEnum): """Update sections""" + ZONE = 0 PREREQ = 1 UPDATE = 2 @@ -46,11 +47,15 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] # ignore the mypy error here as we mean to use a different enum _section_enum = UpdateSection # type: ignore - def __init__(self, zone: Optional[Union[dns.name.Name, str]]=None, - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, - keyring: Optional[Any]=None, keyname: Optional[dns.name.Name]=None, - keyalgorithm: Union[dns.name.Name, str]=dns.tsig.default_algorithm, - id: Optional[int]=None): + def __init__( + self, + zone: Optional[Union[dns.name.Name, str]] = None, + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + keyring: Optional[Any] = None, + keyname: Optional[dns.name.Name] = None, + keyalgorithm: Union[dns.name.Name, str] = dns.tsig.default_algorithm, + id: Optional[int] = None, + ): """Initialize a new DNS Update object. See the documentation of the Message class for a complete @@ -74,8 +79,14 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] rdclass = dns.rdataclass.RdataClass.make(rdclass) self.zone_rdclass = rdclass if self.origin: - self.find_rrset(self.zone, self.origin, rdclass, dns.rdatatype.SOA, - create=True, force_unique=True) + self.find_rrset( + self.zone, + self.origin, + rdclass, + dns.rdatatype.SOA, + create=True, + force_unique=True, + ) if keyring is not None: self.use_tsig(keyring, keyname, algorithm=keyalgorithm) @@ -112,8 +123,9 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] if section is None: section = self.update covers = rd.covers() - rrset = self.find_rrset(section, name, self.zone_rdclass, rd.rdtype, - covers, deleting, True, True) + rrset = self.find_rrset( + section, name, self.zone_rdclass, rd.rdtype, covers, deleting, True, True + ) rrset.add(rd, ttl) def _add(self, replace, section, name, *args): @@ -153,8 +165,7 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] if replace: self.delete(name, rdtype) for s in args: - rd = dns.rdata.from_text(self.zone_rdclass, rdtype, s, - self.origin) + rd = dns.rdata.from_text(self.zone_rdclass, rdtype, s, self.origin) self._add_rr(name, ttl, rd, section=section) def add(self, name: Union[dns.name.Name, str], *args: Any) -> None: @@ -190,9 +201,16 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] if isinstance(name, str): name = dns.name.from_text(name, None) if len(args) == 0: - self.find_rrset(self.update, name, dns.rdataclass.ANY, - dns.rdatatype.ANY, dns.rdatatype.NONE, - dns.rdataclass.ANY, True, True) + self.find_rrset( + self.update, + name, + dns.rdataclass.ANY, + dns.rdatatype.ANY, + dns.rdatatype.NONE, + dns.rdataclass.ANY, + True, + True, + ) elif isinstance(args[0], dns.rdataset.Rdataset): for rds in args: for rd in rds: @@ -205,15 +223,24 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] else: rdtype = dns.rdatatype.RdataType.make(largs.pop(0)) if len(largs) == 0: - self.find_rrset(self.update, name, - self.zone_rdclass, rdtype, - dns.rdatatype.NONE, - dns.rdataclass.ANY, - True, True) + self.find_rrset( + self.update, + name, + self.zone_rdclass, + rdtype, + dns.rdatatype.NONE, + dns.rdataclass.ANY, + True, + True, + ) else: for s in largs: - rd = dns.rdata.from_text(self.zone_rdclass, rdtype, s, # type: ignore[arg-type] - self.origin) + rd = dns.rdata.from_text( + self.zone_rdclass, + rdtype, + s, # type: ignore[arg-type] + self.origin, + ) self._add_rr(name, 0, rd, dns.rdataclass.NONE) def replace(self, name: Union[dns.name.Name, str], *args: Any) -> None: @@ -252,13 +279,21 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] if isinstance(name, str): name = dns.name.from_text(name, None) if len(args) == 0: - self.find_rrset(self.prerequisite, name, - dns.rdataclass.ANY, dns.rdatatype.ANY, - dns.rdatatype.NONE, None, - True, True) - elif isinstance(args[0], dns.rdataset.Rdataset) or \ - isinstance(args[0], dns.rdata.Rdata) or \ - len(args) > 1: + self.find_rrset( + self.prerequisite, + name, + dns.rdataclass.ANY, + dns.rdatatype.ANY, + dns.rdatatype.NONE, + None, + True, + True, + ) + elif ( + isinstance(args[0], dns.rdataset.Rdataset) + or isinstance(args[0], dns.rdata.Rdata) + or len(args) > 1 + ): if not isinstance(args[0], dns.rdataset.Rdataset): # Add a 0 TTL largs = list(args) @@ -268,29 +303,50 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] self._add(False, self.prerequisite, name, *args) else: rdtype = dns.rdatatype.RdataType.make(args[0]) - self.find_rrset(self.prerequisite, name, - dns.rdataclass.ANY, rdtype, - dns.rdatatype.NONE, None, - True, True) - - def absent(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str]=None) -> None: + self.find_rrset( + self.prerequisite, + name, + dns.rdataclass.ANY, + rdtype, + dns.rdatatype.NONE, + None, + True, + True, + ) + + def absent( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str] = None, + ) -> None: """Require that an owner name (and optionally an rdata type) does not exist as a prerequisite to the execution of the update.""" if isinstance(name, str): name = dns.name.from_text(name, None) if rdtype is None: - self.find_rrset(self.prerequisite, name, - dns.rdataclass.NONE, dns.rdatatype.ANY, - dns.rdatatype.NONE, None, - True, True) + self.find_rrset( + self.prerequisite, + name, + dns.rdataclass.NONE, + dns.rdatatype.ANY, + dns.rdatatype.NONE, + None, + True, + True, + ) else: the_rdtype = dns.rdatatype.RdataType.make(rdtype) - self.find_rrset(self.prerequisite, name, - dns.rdataclass.NONE, the_rdtype, - dns.rdatatype.NONE, None, - True, True) + self.find_rrset( + self.prerequisite, + name, + dns.rdataclass.NONE, + the_rdtype, + dns.rdatatype.NONE, + None, + True, + True, + ) def _get_one_rr_per_rrset(self, value): # Updates are always one_rr_per_rrset @@ -300,9 +356,11 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] deleting = None empty = False if section == UpdateSection.ZONE: - if dns.rdataclass.is_metaclass(rdclass) or \ - rdtype != dns.rdatatype.SOA or \ - self.zone: + if ( + dns.rdataclass.is_metaclass(rdclass) + or rdtype != dns.rdatatype.SOA + or self.zone + ): raise dns.exception.FormError else: if not self.zone: @@ -310,10 +368,12 @@ class UpdateMessage(dns.message.Message): # lgtm[py/missing-equals] if rdclass in (dns.rdataclass.ANY, dns.rdataclass.NONE): deleting = rdclass rdclass = self.zone[0].rdclass - empty = (deleting == dns.rdataclass.ANY or - section == UpdateSection.PREREQ) + empty = ( + deleting == dns.rdataclass.ANY or section == UpdateSection.PREREQ + ) return (rdclass, rdtype, deleting, empty) + # backwards compatibility Update = UpdateMessage diff --git a/dns/version.py b/dns/version.py index c6cdf6c..66e8faa 100644 --- a/dns/version.py +++ b/dns/version.py @@ -28,16 +28,31 @@ RELEASELEVEL = 0x00 #: SERIAL SERIAL = 0 -if RELEASELEVEL == 0x0f: # pragma: no cover lgtm[py/unreachable-statement] +if RELEASELEVEL == 0x0F: # pragma: no cover lgtm[py/unreachable-statement] #: version - version = '%d.%d.%d' % (MAJOR, MINOR, MICRO) # lgtm[py/unreachable-statement] + version = "%d.%d.%d" % (MAJOR, MINOR, MICRO) # lgtm[py/unreachable-statement] elif RELEASELEVEL == 0x00: # pragma: no cover lgtm[py/unreachable-statement] - version = '%d.%d.%ddev%d' % (MAJOR, MINOR, MICRO, SERIAL) # lgtm[py/unreachable-statement] -elif RELEASELEVEL == 0x0c: # pragma: no cover lgtm[py/unreachable-statement] - version = '%d.%d.%drc%d' % (MAJOR, MINOR, MICRO, SERIAL) # lgtm[py/unreachable-statement] + version = "%d.%d.%ddev%d" % ( + MAJOR, + MINOR, + MICRO, + SERIAL, + ) # lgtm[py/unreachable-statement] +elif RELEASELEVEL == 0x0C: # pragma: no cover lgtm[py/unreachable-statement] + version = "%d.%d.%drc%d" % ( + MAJOR, + MINOR, + MICRO, + SERIAL, + ) # lgtm[py/unreachable-statement] else: # pragma: no cover lgtm[py/unreachable-statement] - version = '%d.%d.%d%x%d' % (MAJOR, MINOR, MICRO, RELEASELEVEL, SERIAL) # lgtm[py/unreachable-statement] + version = "%d.%d.%d%x%d" % ( + MAJOR, + MINOR, + MICRO, + RELEASELEVEL, + SERIAL, + ) # lgtm[py/unreachable-statement] #: hexversion -hexversion = MAJOR << 24 | MINOR << 16 | MICRO << 8 | RELEASELEVEL << 4 | \ - SERIAL +hexversion = MAJOR << 24 | MINOR << 16 | MICRO << 8 | RELEASELEVEL << 4 | SERIAL diff --git a/dns/versioned.py b/dns/versioned.py index 9ed9cef..5cf29e9 100644 --- a/dns/versioned.py +++ b/dns/versioned.py @@ -5,10 +5,11 @@ from typing import Callable, Deque, Optional, Set, Union import collections + try: import threading as _threading except ImportError: # pragma: no cover - import dummy_threading as _threading # type: ignore + import dummy_threading as _threading # type: ignore import dns.exception import dns.immutable @@ -36,15 +37,25 @@ Transaction = dns.zone.Transaction class Zone(dns.zone.Zone): # lgtm[py/missing-equals] - __slots__ = ['_versions', '_versions_lock', '_write_txn', - '_write_waiters', '_write_event', '_pruning_policy', - '_readers'] + __slots__ = [ + "_versions", + "_versions_lock", + "_write_txn", + "_write_waiters", + "_write_event", + "_pruning_policy", + "_readers", + ] node_factory = Node - def __init__(self, origin: Optional[Union[dns.name.Name, str]], - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, relativize: bool=True, - pruning_policy: Optional[Callable[['Zone', Version], Optional[bool]]]=None): + def __init__( + self, + origin: Optional[Union[dns.name.Name, str]], + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + relativize: bool = True, + pruning_policy: Optional[Callable[["Zone", Version], Optional[bool]]] = None, + ): """Initialize a versioned zone object. *origin* is the origin of the zone. It may be a ``dns.name.Name``, @@ -71,13 +82,15 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] self._write_event: Optional[_threading.Event] = None self._write_waiters: Deque[_threading.Event] = collections.deque() self._readers: Set[Transaction] = set() - self._commit_version_unlocked(None, - WritableVersion(self, replacement=True), - origin) + self._commit_version_unlocked( + None, WritableVersion(self, replacement=True), origin + ) - def reader(self, id: Optional[int]=None, serial: Optional[int]=None) -> Transaction: # pylint: disable=arguments-differ + def reader( + self, id: Optional[int] = None, serial: Optional[int] = None + ) -> Transaction: # pylint: disable=arguments-differ if id is not None and serial is not None: - raise ValueError('cannot specify both id and serial') + raise ValueError("cannot specify both id and serial") with self._version_lock: if id is not None: version = None @@ -86,7 +99,7 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] version = v break if version is None: - raise KeyError('version not found') + raise KeyError("version not found") elif serial is not None: if self.relativize: oname = dns.name.empty @@ -102,14 +115,14 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] version = v break if version is None: - raise KeyError('serial not found') + raise KeyError("serial not found") else: version = self._versions[-1] txn = Transaction(self, False, version) self._readers.add(txn) return txn - def writer(self, replacement: bool=False) -> Transaction: + def writer(self, replacement: bool = False) -> Transaction: event = None while True: with self._version_lock: @@ -123,8 +136,9 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] # give up the lock, so that we hold the lock as # short a time as possible. This is why we call # _setup_version() below. - self._write_txn = Transaction(self, replacement, - make_immutable=True) + self._write_txn = Transaction( + self, replacement, make_immutable=True + ) # give up our exclusive right to make a Transaction self._write_event = None break @@ -165,6 +179,7 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] # pylint: disable=unused-argument def _default_pruning_policy(self, zone, version): return True + # pylint: enable=unused-argument def _prune_versions_unlocked(self): @@ -180,8 +195,9 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] least_kept = min(txn.version.id for txn in self._readers) else: least_kept = self._versions[-1].id - while self._versions[0].id < least_kept and \ - self._pruning_policy(self, self._versions[0]): + while self._versions[0].id < least_kept and self._pruning_policy( + self, self._versions[0] + ): self._versions.popleft() def set_max_versions(self, max_versions: Optional[int]) -> None: @@ -189,16 +205,22 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] of versions """ if max_versions is not None and max_versions < 1: - raise ValueError('max versions must be at least 1') + raise ValueError("max versions must be at least 1") if max_versions is None: + def policy(zone, _): # pylint: disable=unused-argument return False + else: + def policy(zone, _): return len(zone._versions) > max_versions + self.set_pruning_policy(policy) - def set_pruning_policy(self, policy: Optional[Callable[['Zone', Version], Optional[bool]]]) -> None: + def set_pruning_policy( + self, policy: Optional[Callable[["Zone", Version], Optional[bool]]] + ) -> None: """Set the pruning policy for the zone. The *policy* function takes a `Version` and returns `True` if @@ -251,7 +273,9 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] id = 1 return id - def find_node(self, name: Union[dns.name.Name, str], create: bool=False) -> dns.node.Node: + def find_node( + self, name: Union[dns.name.Name, str], create: bool = False + ) -> dns.node.Node: if create: raise UseTransaction return super().find_node(name) @@ -259,19 +283,25 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] def delete_node(self, name: Union[dns.name.Name, str]) -> None: raise UseTransaction - def find_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE, - create: bool=False) -> dns.rdataset.Rdataset: + def find_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + create: bool = False, + ) -> dns.rdataset.Rdataset: if create: raise UseTransaction rdataset = super().find_rdataset(name, rdtype, covers) return dns.rdataset.ImmutableRdataset(rdataset) - def get_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE, - create: bool=False) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + create: bool = False, + ) -> Optional[dns.rdataset.Rdataset]: if create: raise UseTransaction rdataset = super().get_rdataset(name, rdtype, covers) @@ -280,10 +310,15 @@ class Zone(dns.zone.Zone): # lgtm[py/missing-equals] else: return None - def delete_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) -> None: + def delete_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> None: raise UseTransaction - def replace_rdataset(self, name: Union[dns.name.Name, str], replacement: dns.rdataset.Rdataset) -> None: + def replace_rdataset( + self, name: Union[dns.name.Name, str], replacement: dns.rdataset.Rdataset + ) -> None: raise UseTransaction diff --git a/dns/win32util.py b/dns/win32util.py index f4ded20..7a17b0b 100644 --- a/dns/win32util.py +++ b/dns/win32util.py @@ -1,6 +1,6 @@ import sys -if sys.platform == 'win32': +if sys.platform == "win32": from typing import Any @@ -14,9 +14,10 @@ if sys.platform == 'win32': try: import threading as _threading except ImportError: # pragma: no cover - import dummy_threading as _threading # type: ignore + import dummy_threading as _threading # type: ignore import pythoncom import wmi + _have_wmi = True except Exception: _have_wmi = False @@ -25,7 +26,7 @@ if sys.platform == 'win32': # Sometimes DHCP servers add a '.' prefix to the default domain, and # Windows just stores such values in the registry (see #687). # Check for this and fix it. - if domain.startswith('.'): + if domain.startswith("."): domain = domain[1:] return dns.name.from_text(domain) @@ -36,6 +37,7 @@ if sys.platform == 'win32': self.search = [] if _have_wmi: + class _WMIGetter(_threading.Thread): def __init__(self): super().__init__() @@ -49,8 +51,10 @@ if sys.platform == 'win32': if interface.IPEnabled: self.info.domain = _config_domain(interface.DNSDomain) self.info.nameservers = list(interface.DNSServerSearchOrder) - self.info.search = [dns.name.from_text(x) for x in - interface.DNSDomainSuffixSearchOrder] + self.info.search = [ + dns.name.from_text(x) + for x in interface.DNSDomainSuffixSearchOrder + ] break finally: pythoncom.CoUninitialize() @@ -61,11 +65,12 @@ if sys.platform == 'win32': self.start() self.join() return self.info + else: + class _WMIGetter: # type: ignore pass - class _RegistryGetter: def __init__(self): self.info = DnsInfo() @@ -76,13 +81,13 @@ if sys.platform == 'win32': # delimiter in between ' ' and ',' (and vice-versa) in various # versions of windows. # - if entry.find(' ') >= 0: - split_char = ' ' - elif entry.find(',') >= 0: - split_char = ',' + if entry.find(" ") >= 0: + split_char = " " + elif entry.find(",") >= 0: + split_char = "," else: # probably a singleton; treat as a space-separated list. - split_char = ' ' + split_char = " " return split_char def _config_nameservers(self, nameservers): @@ -102,38 +107,38 @@ if sys.platform == 'win32': def _config_fromkey(self, key, always_try_domain): try: - servers, _ = winreg.QueryValueEx(key, 'NameServer') + servers, _ = winreg.QueryValueEx(key, "NameServer") except WindowsError: servers = None if servers: self._config_nameservers(servers) if servers or always_try_domain: try: - dom, _ = winreg.QueryValueEx(key, 'Domain') + dom, _ = winreg.QueryValueEx(key, "Domain") if dom: self.info.domain = _config_domain(dom) except WindowsError: pass else: try: - servers, _ = winreg.QueryValueEx(key, 'DhcpNameServer') + servers, _ = winreg.QueryValueEx(key, "DhcpNameServer") except WindowsError: servers = None if servers: self._config_nameservers(servers) try: - dom, _ = winreg.QueryValueEx(key, 'DhcpDomain') + dom, _ = winreg.QueryValueEx(key, "DhcpDomain") if dom: self.info.domain = _config_domain(dom) except WindowsError: pass try: - search, _ = winreg.QueryValueEx(key, 'SearchList') + search, _ = winreg.QueryValueEx(key, "SearchList") except WindowsError: search = None if search is None: try: - search, _ = winreg.QueryValueEx(key, 'DhcpSearchList') + search, _ = winreg.QueryValueEx(key, "DhcpSearchList") except WindowsError: search = None if search: @@ -150,25 +155,27 @@ if sys.platform == 'win32': # from Windows 2000 through Vista. connection_key = winreg.OpenKey( lm, - r'SYSTEM\CurrentControlSet\Control\Network' - r'\{4D36E972-E325-11CE-BFC1-08002BE10318}' - r'\%s\Connection' % guid) + r"SYSTEM\CurrentControlSet\Control\Network" + r"\{4D36E972-E325-11CE-BFC1-08002BE10318}" + r"\%s\Connection" % guid, + ) try: # The PnpInstanceID points to a key inside Enum (pnp_id, ttype) = winreg.QueryValueEx( - connection_key, 'PnpInstanceID') + connection_key, "PnpInstanceID" + ) if ttype != winreg.REG_SZ: raise ValueError # pragma: no cover device_key = winreg.OpenKey( - lm, r'SYSTEM\CurrentControlSet\Enum\%s' % pnp_id) + lm, r"SYSTEM\CurrentControlSet\Enum\%s" % pnp_id + ) try: # Get ConfigFlags for this device - (flags, ttype) = winreg.QueryValueEx( - device_key, 'ConfigFlags') + (flags, ttype) = winreg.QueryValueEx(device_key, "ConfigFlags") if ttype != winreg.REG_DWORD: raise ValueError # pragma: no cover @@ -194,17 +201,19 @@ if sys.platform == 'win32': lm = winreg.ConnectRegistry(None, winreg.HKEY_LOCAL_MACHINE) try: - tcp_params = winreg.OpenKey(lm, - r'SYSTEM\CurrentControlSet' - r'\Services\Tcpip\Parameters') + tcp_params = winreg.OpenKey( + lm, r"SYSTEM\CurrentControlSet" r"\Services\Tcpip\Parameters" + ) try: self._config_fromkey(tcp_params, True) finally: tcp_params.Close() - interfaces = winreg.OpenKey(lm, - r'SYSTEM\CurrentControlSet' - r'\Services\Tcpip\Parameters' - r'\Interfaces') + interfaces = winreg.OpenKey( + lm, + r"SYSTEM\CurrentControlSet" + r"\Services\Tcpip\Parameters" + r"\Interfaces", + ) try: i = 0 while True: diff --git a/dns/wire.py b/dns/wire.py index 905930f..cadf168 100644 --- a/dns/wire.py +++ b/dns/wire.py @@ -8,8 +8,9 @@ import struct import dns.exception import dns.name + class Parser: - def __init__(self, wire: bytes, current: int=0): + def __init__(self, wire: bytes, current: int = 0): self.wire = wire self.current = 0 self.end = len(self.wire) @@ -24,34 +25,34 @@ class Parser: assert size >= 0 if size > self.remaining(): raise dns.exception.FormError - output = self.wire[self.current:self.current + size] + output = self.wire[self.current : self.current + size] self.current += size self.furthest = max(self.furthest, self.current) return output - def get_counted_bytes(self, length_size: int=1) -> bytes: - length = int.from_bytes(self.get_bytes(length_size), 'big') + def get_counted_bytes(self, length_size: int = 1) -> bytes: + length = int.from_bytes(self.get_bytes(length_size), "big") return self.get_bytes(length) def get_remaining(self) -> bytes: return self.get_bytes(self.remaining()) def get_uint8(self) -> int: - return struct.unpack('!B', self.get_bytes(1))[0] + return struct.unpack("!B", self.get_bytes(1))[0] def get_uint16(self) -> int: - return struct.unpack('!H', self.get_bytes(2))[0] + return struct.unpack("!H", self.get_bytes(2))[0] def get_uint32(self) -> int: - return struct.unpack('!I', self.get_bytes(4))[0] + return struct.unpack("!I", self.get_bytes(4))[0] def get_uint48(self) -> int: - return int.from_bytes(self.get_bytes(6), 'big') + return int.from_bytes(self.get_bytes(6), "big") def get_struct(self, format: str) -> Tuple: return struct.unpack(format, self.get_bytes(struct.calcsize(format))) - def get_name(self, origin: Optional['dns.name.Name']=None) -> 'dns.name.Name': + def get_name(self, origin: Optional["dns.name.Name"] = None) -> "dns.name.Name": name = dns.name.from_wire_parser(self) if origin: name = name.relativize(origin) @@ -33,7 +33,7 @@ class TransferError(dns.exception.DNSException): """A zone transfer response got a non-zero rcode.""" def __init__(self, rcode): - message = 'Zone transfer error: %s' % dns.rcode.to_text(rcode) + message = "Zone transfer error: %s" % dns.rcode.to_text(rcode) super().__init__(message) self.rcode = rcode @@ -51,9 +51,13 @@ class Inbound: State machine for zone transfers. """ - def __init__(self, txn_manager: dns.transaction.TransactionManager, - rdtype: dns.rdatatype.RdataType=dns.rdatatype.AXFR, - serial: Optional[int]=None, is_udp: bool=False): + def __init__( + self, + txn_manager: dns.transaction.TransactionManager, + rdtype: dns.rdatatype.RdataType = dns.rdatatype.AXFR, + serial: Optional[int] = None, + is_udp: bool = False, + ): """Initialize an inbound zone transfer. *txn_manager* is a :py:class:`dns.transaction.TransactionManager`. @@ -71,9 +75,9 @@ class Inbound: self.rdtype = rdtype if rdtype == dns.rdatatype.IXFR: if serial is None: - raise ValueError('a starting serial must be supplied for IXFRs') + raise ValueError("a starting serial must be supplied for IXFRs") elif is_udp: - raise ValueError('is_udp specified for AXFR') + raise ValueError("is_udp specified for AXFR") self.serial = serial self.is_udp = is_udp (_, _, self.origin) = txn_manager.origin_information() @@ -113,8 +117,9 @@ class Inbound: # the origin. # if not message.answer or message.answer[0].name != self.origin: - raise dns.exception.FormError("No answer or RRset not " - "for zone origin") + raise dns.exception.FormError( + "No answer or RRset not " "for zone origin" + ) rrset = message.answer[0] rdataset = rrset if rdataset.rdtype != dns.rdatatype.SOA: @@ -127,8 +132,7 @@ class Inbound: # We're already up-to-date. # self.done = True - elif dns.serial.Serial(self.soa_rdataset[0].serial) < \ - self.serial: + elif dns.serial.Serial(self.soa_rdataset[0].serial) < self.serial: # It went backwards! raise SerialWentBackwards else: @@ -153,8 +157,7 @@ class Inbound: if self.done: raise dns.exception.FormError("answers after final SOA") assert self.txn is not None # for mypy - if rdataset.rdtype == dns.rdatatype.SOA and \ - name == self.origin: + if rdataset.rdtype == dns.rdatatype.SOA and name == self.origin: # # Every time we see an origin SOA delete_mode inverts # @@ -166,20 +169,23 @@ class Inbound: # check that we're seeing the record in the expected # part of the response. # - if rdataset == self.soa_rdataset and \ - (self.rdtype == dns.rdatatype.AXFR or - (self.rdtype == dns.rdatatype.IXFR and - self.delete_mode)): + if rdataset == self.soa_rdataset and ( + self.rdtype == dns.rdatatype.AXFR + or (self.rdtype == dns.rdatatype.IXFR and self.delete_mode) + ): # # This is the final SOA # if self.expecting_SOA: # We got an empty IXFR sequence! - raise dns.exception.FormError('empty IXFR sequence') - if self.rdtype == dns.rdatatype.IXFR \ - and self.serial != rdataset[0].serial: - raise dns.exception.FormError('unexpected end of IXFR ' - 'sequence') + raise dns.exception.FormError("empty IXFR sequence") + if ( + self.rdtype == dns.rdatatype.IXFR + and self.serial != rdataset[0].serial + ): + raise dns.exception.FormError( + "unexpected end of IXFR " "sequence" + ) self.txn.replace(name, rdataset) self.txn.commit() self.txn = None @@ -194,15 +200,17 @@ class Inbound: # This is the start of an IXFR deletion set if rdataset[0].serial != self.serial: raise dns.exception.FormError( - "IXFR base serial mismatch") + "IXFR base serial mismatch" + ) else: # This is the start of an IXFR addition set self.serial = rdataset[0].serial self.txn.replace(name, rdataset) else: # We saw a non-final SOA for the origin in an AXFR. - raise dns.exception.FormError('unexpected origin SOA ' - 'in AXFR') + raise dns.exception.FormError( + "unexpected origin SOA " "in AXFR" + ) continue if self.expecting_SOA: # @@ -229,7 +237,7 @@ class Inbound: # This is a UDP IXFR and we didn't get to done, and we didn't # get the proper "truncated" response # - raise dns.exception.FormError('unexpected end of UDP IXFR') + raise dns.exception.FormError("unexpected end of UDP IXFR") return self.done # @@ -245,12 +253,18 @@ class Inbound: return False -def make_query(txn_manager: dns.transaction.TransactionManager, serial: Optional[int]=0, - use_edns: Optional[Union[int, bool]]=None, ednsflags: Optional[int]=None, payload: Optional[int]=None, - request_payload: Optional[int]=None, options: Optional[List[dns.edns.Option]]=None, - keyring: Any=None, keyname: Optional[dns.name.Name]=None, - keyalgorithm: Union[dns.name.Name, str]=dns.tsig.default_algorithm) \ - -> Tuple[dns.message.QueryMessage, Optional[int]]: +def make_query( + txn_manager: dns.transaction.TransactionManager, + serial: Optional[int] = 0, + use_edns: Optional[Union[int, bool]] = None, + ednsflags: Optional[int] = None, + payload: Optional[int] = None, + request_payload: Optional[int] = None, + options: Optional[List[dns.edns.Option]] = None, + keyring: Any = None, + keyname: Optional[dns.name.Name] = None, + keyalgorithm: Union[dns.name.Name, str] = dns.tsig.default_algorithm, +) -> Tuple[dns.message.QueryMessage, Optional[int]]: """Make an AXFR or IXFR query. *txn_manager* is a ``dns.transaction.TransactionManager``, typically a @@ -272,14 +286,14 @@ def make_query(txn_manager: dns.transaction.TransactionManager, serial: Optional """ (zone_origin, _, origin) = txn_manager.origin_information() if zone_origin is None: - raise ValueError('no zone origin') + raise ValueError("no zone origin") if serial is None: rdtype = dns.rdatatype.AXFR elif not isinstance(serial, int): - raise ValueError('serial is not an integer') + raise ValueError("serial is not an integer") elif serial == 0: with txn_manager.reader() as txn: - rdataset = txn.get(origin, 'SOA') + rdataset = txn.get(origin, "SOA") if rdataset: serial = rdataset[0].serial rdtype = dns.rdatatype.IXFR @@ -289,20 +303,30 @@ def make_query(txn_manager: dns.transaction.TransactionManager, serial: Optional elif serial > 0 and serial < 4294967296: rdtype = dns.rdatatype.IXFR else: - raise ValueError('serial out-of-range') + raise ValueError("serial out-of-range") rdclass = txn_manager.get_class() - q = dns.message.make_query(zone_origin, rdtype, rdclass, - use_edns, False, ednsflags, payload, - request_payload, options) + q = dns.message.make_query( + zone_origin, + rdtype, + rdclass, + use_edns, + False, + ednsflags, + payload, + request_payload, + options, + ) if serial is not None: - rdata = dns.rdata.from_text(rdclass, 'SOA', f'. . {serial} 0 0 0 0') - rrset = q.find_rrset(q.authority, zone_origin, rdclass, - dns.rdatatype.SOA, create=True) + rdata = dns.rdata.from_text(rdclass, "SOA", f". . {serial} 0 0 0 0") + rrset = q.find_rrset( + q.authority, zone_origin, rdclass, dns.rdatatype.SOA, create=True + ) rrset.add(rdata, 0) if keyring is not None: q.use_tsig(keyring, keyname, algorithm=keyalgorithm) return (q, serial) + def extract_serial_from_query(query: dns.message.Message) -> Optional[int]: """Extract the SOA serial number from query if it is an IXFR and return it, otherwise return None. @@ -313,12 +337,13 @@ def extract_serial_from_query(query: dns.message.Message) -> Optional[int]: an appropriate SOA RRset in the authority section. """ if not isinstance(query, dns.message.QueryMessage): - raise ValueError('query not a QueryMessage') + raise ValueError("query not a QueryMessage") question = query.question[0] if question.rdtype == dns.rdatatype.AXFR: return None elif question.rdtype != dns.rdatatype.IXFR: raise ValueError("query is not an AXFR or IXFR") - soa = query.find_rrset(query.authority, question.name, question.rdclass, - dns.rdatatype.SOA) + soa = query.find_rrset( + query.authority, question.name, question.rdclass, dns.rdatatype.SOA + ) return soa[0].serial diff --git a/dns/zone.py b/dns/zone.py index d57838c..d0d9928 100644 --- a/dns/zone.py +++ b/dns/zone.py @@ -97,10 +97,14 @@ class Zone(dns.transaction.TransactionManager): node_factory = dns.node.Node - __slots__ = ['rdclass', 'origin', 'nodes', 'relativize'] - - def __init__(self, origin: Optional[Union[dns.name.Name, str]], - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, relativize: bool=True): + __slots__ = ["rdclass", "origin", "nodes", "relativize"] + + def __init__( + self, + origin: Optional[Union[dns.name.Name, str]], + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + relativize: bool = True, + ): """Initialize a zone object. *origin* is the origin of the zone. It may be a ``dns.name.Name``, @@ -117,8 +121,9 @@ class Zone(dns.transaction.TransactionManager): if isinstance(origin, str): origin = dns.name.from_text(origin) elif not isinstance(origin, dns.name.Name): - raise ValueError("origin parameter must be convertible to a " - "DNS name") + raise ValueError( + "origin parameter must be convertible to a " "DNS name" + ) if not origin.is_absolute(): raise ValueError("origin parameter must be an absolute name") self.origin = origin @@ -135,9 +140,11 @@ class Zone(dns.transaction.TransactionManager): if not isinstance(other, Zone): return False - if self.rdclass != other.rdclass or \ - self.origin != other.origin or \ - self.nodes != other.nodes: + if ( + self.rdclass != other.rdclass + or self.origin != other.origin + or self.nodes != other.nodes + ): return False return True @@ -159,16 +166,15 @@ class Zone(dns.transaction.TransactionManager): # This should probably never happen as other code (e.g. # _rr_line) will notice the lack of an origin before us, but # we check just in case! - raise KeyError('no zone origin is defined') + raise KeyError("no zone origin is defined") if not name.is_subdomain(self.origin): - raise KeyError( - "name parameter must be a subdomain of the zone origin") + raise KeyError("name parameter must be a subdomain of the zone origin") if self.relativize: name = name.relativize(self.origin) elif not self.relativize: # We have a relative name in a non-relative zone, so derelativize. if self.origin is None: - raise KeyError('no zone origin is defined') + raise KeyError("no zone origin is defined") name = name.derelativize(self.origin) return name @@ -204,7 +210,9 @@ class Zone(dns.transaction.TransactionManager): key = self._validate_name(key) return key in self.nodes - def find_node(self, name: Union[dns.name.Name, str], create: bool=False) -> dns.node.Node: + def find_node( + self, name: Union[dns.name.Name, str], create: bool = False + ) -> dns.node.Node: """Find a node in the zone, possibly creating it. *name*: the name of the node to find. @@ -230,7 +238,9 @@ class Zone(dns.transaction.TransactionManager): self.nodes[name] = node return node - def get_node(self, name: Union[dns.name.Name, str], create: bool=False) -> Optional[dns.node.Node]: + def get_node( + self, name: Union[dns.name.Name, str], create: bool = False + ) -> Optional[dns.node.Node]: """Get a node in the zone, possibly creating it. This method is like ``find_node()``, except it returns None instead @@ -272,10 +282,13 @@ class Zone(dns.transaction.TransactionManager): if name in self.nodes: del self.nodes[name] - def find_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE, - create: bool=False) -> dns.rdataset.Rdataset: + def find_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + create: bool = False, + ) -> dns.rdataset.Rdataset: """Look for an rdataset with the specified name and type in the zone, and return an rdataset encapsulating it. @@ -316,10 +329,13 @@ class Zone(dns.transaction.TransactionManager): node = self.find_node(the_name, create) return node.find_rdataset(self.rdclass, the_rdtype, the_covers, create) - def get_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE, - create: bool=False) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + create: bool = False, + ) -> Optional[dns.rdataset.Rdataset]: """Look for an rdataset with the specified name and type in the zone. This method is like ``find_rdataset()``, except it returns None instead @@ -361,34 +377,33 @@ class Zone(dns.transaction.TransactionManager): rdataset = None return rdataset - def delete_rdataset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) -> None: + def delete_rdataset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> None: """Delete the rdataset matching *rdtype* and *covers*, if it exists at the node specified by *name*. - It is not an error if the node does not exist, or if there is no - matching rdataset at the node. + It is not an error if the node does not exist, or if there is no matching + rdataset at the node. - If the node has no rdatasets after the deletion, it will itself - be deleted. + If the node has no rdatasets after the deletion, it will itself be deleted. - *name*: the name of the node to find. - The value may be a ``dns.name.Name`` or a ``str``. If absolute, the - name must be a subdomain of the zone's origin. If ``zone.relativize`` - is ``True``, then the name will be relativized. + *name*: the name of the node to find. The value may be a ``dns.name.Name`` or a + ``str``. If absolute, the name must be a subdomain of the zone's origin. If + ``zone.relativize`` is ``True``, then the name will be relativized. *rdtype*, a ``dns.rdatatype.RdataType`` or ``str``, the rdata type desired. - *covers*, a ``dns.rdatatype.RdataType`` or ``str`` or ``None``, the covered type. - Usually this value is ``dns.rdatatype.NONE``, but if the - rdtype is ``dns.rdatatype.SIG`` or ``dns.rdatatype.RRSIG``, - then the covers value will be the rdata type the SIG/RRSIG - covers. The library treats the SIG and RRSIG types as if they - were a family of types, e.g. RRSIG(A), RRSIG(NS), RRSIG(SOA). - This makes RRSIGs much easier to work with than if RRSIGs - covering different rdata types were aggregated into a single - RRSIG rdataset. + *covers*, a ``dns.rdatatype.RdataType`` or ``str`` or ``None``, the covered + type. Usually this value is ``dns.rdatatype.NONE``, but if the rdtype is + ``dns.rdatatype.SIG`` or ``dns.rdatatype.RRSIG``, then the covers value will be + the rdata type the SIG/RRSIG covers. The library treats the SIG and RRSIG types + as if they were a family of types, e.g. RRSIG(A), RRSIG(NS), RRSIG(SOA). This + makes RRSIGs much easier to work with than if RRSIGs covering different rdata + types were aggregated into a single RRSIG rdataset. """ the_name = self._validate_name(name) @@ -400,8 +415,9 @@ class Zone(dns.transaction.TransactionManager): if len(node) == 0: self.delete_node(the_name) - def replace_rdataset(self, name: Union[dns.name.Name, str], - replacement: dns.rdataset.Rdataset) -> None: + def replace_rdataset( + self, name: Union[dns.name.Name, str], replacement: dns.rdataset.Rdataset + ) -> None: """Replace an rdataset at name. It is not an error if there is no rdataset matching I{replacement}. @@ -421,13 +437,16 @@ class Zone(dns.transaction.TransactionManager): """ if replacement.rdclass != self.rdclass: - raise ValueError('replacement.rdclass != zone.rdclass') + raise ValueError("replacement.rdclass != zone.rdclass") node = self.find_node(name, True) node.replace_rdataset(replacement) - def find_rrset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) -> dns.rrset.RRset: + def find_rrset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> dns.rrset.RRset: """Look for an rdataset with the specified name and type in the zone, and return an RRset encapsulating it. @@ -474,9 +493,12 @@ class Zone(dns.transaction.TransactionManager): rrset.update(rdataset) return rrset - def get_rrset(self, name: Union[dns.name.Name, str], - rdtype: Union[dns.rdatatype.RdataType, str], - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) -> Optional[dns.rrset.RRset]: + def get_rrset( + self, + name: Union[dns.name.Name, str], + rdtype: Union[dns.rdatatype.RdataType, str], + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> Optional[dns.rrset.RRset]: """Look for an rdataset with the specified name and type in the zone, and return an RRset encapsulating it. @@ -520,9 +542,11 @@ class Zone(dns.transaction.TransactionManager): rrset = None return rrset - def iterate_rdatasets(self, rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.ANY, - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) \ - -> Iterator[Tuple[dns.name.Name, dns.rdataset.Rdataset]]: + def iterate_rdatasets( + self, + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.ANY, + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> Iterator[Tuple[dns.name.Name, dns.rdataset.Rdataset]]: """Return a generator which yields (name, rdataset) tuples for all rdatasets in the zone which have the specified *rdtype* and *covers*. If *rdtype* is ``dns.rdatatype.ANY``, the default, @@ -545,13 +569,16 @@ class Zone(dns.transaction.TransactionManager): covers = dns.rdatatype.RdataType.make(covers) for (name, node) in self.items(): for rds in node: - if rdtype == dns.rdatatype.ANY or \ - (rds.rdtype == rdtype and rds.covers == covers): + if rdtype == dns.rdatatype.ANY or ( + rds.rdtype == rdtype and rds.covers == covers + ): yield (name, rds) - def iterate_rdatas(self, rdtype: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.ANY, - covers: Union[dns.rdatatype.RdataType, str]=dns.rdatatype.NONE) \ - -> Iterator[Tuple[dns.name.Name, int, dns.rdata.Rdata]]: + def iterate_rdatas( + self, + rdtype: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.ANY, + covers: Union[dns.rdatatype.RdataType, str] = dns.rdatatype.NONE, + ) -> Iterator[Tuple[dns.name.Name, int, dns.rdata.Rdata]]: """Return a generator which yields (name, ttl, rdata) tuples for all rdatas in the zone which have the specified *rdtype* and *covers*. If *rdtype* is ``dns.rdatatype.ANY``, the default, @@ -574,13 +601,21 @@ class Zone(dns.transaction.TransactionManager): covers = dns.rdatatype.RdataType.make(covers) for (name, node) in self.items(): for rds in node: - if rdtype == dns.rdatatype.ANY or \ - (rds.rdtype == rdtype and rds.covers == covers): + if rdtype == dns.rdatatype.ANY or ( + rds.rdtype == rdtype and rds.covers == covers + ): for rdata in rds: yield (name, rds.ttl, rdata) - def to_file(self, f: Any, sorted: bool=True, relativize: bool=True, nl: Optional[str]=None, - want_comments: bool=False, want_origin: bool=False) -> None: + def to_file( + self, + f: Any, + sorted: bool = True, + relativize: bool = True, + nl: Optional[str] = None, + want_comments: bool = False, + want_origin: bool = False, + ) -> None: """Write a zone to a file. *f*, a file or `str`. If *f* is a string, it is treated @@ -610,18 +645,18 @@ class Zone(dns.transaction.TransactionManager): with contextlib.ExitStack() as stack: if isinstance(f, str): - f = stack.enter_context(open(f, 'wb')) + f = stack.enter_context(open(f, "wb")) # must be in this way, f.encoding may contain None, or even # attribute may not be there - file_enc = getattr(f, 'encoding', None) + file_enc = getattr(f, "encoding", None) if file_enc is None: - file_enc = 'utf-8' + file_enc = "utf-8" if nl is None: # binary mode, '\n' is not enough nl_b = os.linesep.encode(file_enc) - nl = '\n' + nl = "\n" elif isinstance(nl, str): nl_b = nl.encode(file_enc) else: @@ -630,7 +665,7 @@ class Zone(dns.transaction.TransactionManager): if want_origin: assert self.origin is not None - l = '$ORIGIN ' + self.origin.to_text() + l = "$ORIGIN " + self.origin.to_text() l_b = l.encode(file_enc) try: f.write(l_b) @@ -645,9 +680,12 @@ class Zone(dns.transaction.TransactionManager): else: names = self.keys() for n in names: - l = self[n].to_text(n, origin=self.origin, - relativize=relativize, - want_comments=want_comments) + l = self[n].to_text( + n, + origin=self.origin, + relativize=relativize, + want_comments=want_comments, + ) l_b = l.encode(file_enc) try: @@ -657,8 +695,14 @@ class Zone(dns.transaction.TransactionManager): f.write(l) f.write(nl) - def to_text(self, sorted: bool=True, relativize: bool=True, nl: Optional[str]=None, - want_comments: bool=False, want_origin: bool=False) -> str: + def to_text( + self, + sorted: bool = True, + relativize: bool = True, + nl: Optional[str] = None, + want_comments: bool = False, + want_origin: bool = False, + ) -> str: """Return a zone's text as though it were written to a file. *sorted*, a ``bool``. If True, the default, then the file @@ -685,8 +729,7 @@ class Zone(dns.transaction.TransactionManager): Returns a ``str``. """ temp_buffer = io.StringIO() - self.to_file(temp_buffer, sorted, relativize, nl, want_comments, - want_origin) + self.to_file(temp_buffer, sorted, relativize, nl, want_comments, want_origin) return_value = temp_buffer.getvalue() temp_buffer.close() return return_value @@ -710,7 +753,9 @@ class Zone(dns.transaction.TransactionManager): if self.get_rdataset(name, dns.rdatatype.NS) is None: raise NoNS - def get_soa(self, txn: Optional[dns.transaction.Transaction]=None) -> dns.rdtypes.ANY.SOA.SOA: + def get_soa( + self, txn: Optional[dns.transaction.Transaction] = None + ) -> dns.rdtypes.ANY.SOA.SOA: """Get the zone SOA rdata. Raises ``dns.zone.NoSOA`` if there is no SOA RRset. @@ -734,7 +779,11 @@ class Zone(dns.transaction.TransactionManager): raise NoSOA return soa[0] - def _compute_digest(self, hash_algorithm: DigestHashAlgorithm, scheme: DigestScheme=DigestScheme.SIMPLE) -> bytes: + def _compute_digest( + self, + hash_algorithm: DigestHashAlgorithm, + scheme: DigestScheme = DigestScheme.SIMPLE, + ) -> bytes: hashinfo = _digest_hashers.get(hash_algorithm) if not hashinfo: raise UnsupportedDigestHashAlgorithm @@ -749,30 +798,35 @@ class Zone(dns.transaction.TransactionManager): hasher = hashinfo() for (name, node) in sorted(self.items()): rrnamebuf = name.to_digestable(self.origin) - for rdataset in sorted(node, - key=lambda rds: (rds.rdtype, rds.covers)): - if name == origin_name and \ - dns.rdatatype.ZONEMD in (rdataset.rdtype, rdataset.covers): + for rdataset in sorted(node, key=lambda rds: (rds.rdtype, rds.covers)): + if name == origin_name and dns.rdatatype.ZONEMD in ( + rdataset.rdtype, + rdataset.covers, + ): continue - rrfixed = struct.pack('!HHI', rdataset.rdtype, - rdataset.rdclass, rdataset.ttl) - rdatas = [rdata.to_digestable(self.origin) - for rdata in rdataset] + rrfixed = struct.pack( + "!HHI", rdataset.rdtype, rdataset.rdclass, rdataset.ttl + ) + rdatas = [rdata.to_digestable(self.origin) for rdata in rdataset] for rdata in sorted(rdatas): - rrlen = struct.pack('!H', len(rdata)) + rrlen = struct.pack("!H", len(rdata)) hasher.update(rrnamebuf + rrfixed + rrlen + rdata) return hasher.digest() - def compute_digest(self, hash_algorithm: DigestHashAlgorithm, - scheme: DigestScheme=DigestScheme.SIMPLE) -> dns.rdtypes.ANY.ZONEMD.ZONEMD: + def compute_digest( + self, + hash_algorithm: DigestHashAlgorithm, + scheme: DigestScheme = DigestScheme.SIMPLE, + ) -> dns.rdtypes.ANY.ZONEMD.ZONEMD: serial = self.get_soa().serial digest = self._compute_digest(hash_algorithm, scheme) - return dns.rdtypes.ANY.ZONEMD.ZONEMD(self.rdclass, - dns.rdatatype.ZONEMD, - serial, scheme, hash_algorithm, - digest) + return dns.rdtypes.ANY.ZONEMD.ZONEMD( + self.rdclass, dns.rdatatype.ZONEMD, serial, scheme, hash_algorithm, digest + ) - def verify_digest(self, zonemd: Optional[dns.rdtypes.ANY.ZONEMD.ZONEMD]=None) -> None: + def verify_digest( + self, zonemd: Optional[dns.rdtypes.ANY.ZONEMD.ZONEMD] = None + ) -> None: digests: Union[dns.rdataset.Rdataset, List[dns.rdtypes.ANY.ZONEMD.ZONEMD]] if zonemd: digests = [zonemd] @@ -784,8 +838,7 @@ class Zone(dns.transaction.TransactionManager): digests = rds for digest in digests: try: - computed = self._compute_digest(digest.hash_algorithm, - digest.scheme) + computed = self._compute_digest(digest.hash_algorithm, digest.scheme) if computed == digest.digest: return except Exception: @@ -794,16 +847,17 @@ class Zone(dns.transaction.TransactionManager): # TransactionManager methods - def reader(self) -> 'Transaction': - return Transaction(self, False, - Version(self, 1, self.nodes, self.origin)) + def reader(self) -> "Transaction": + return Transaction(self, False, Version(self, 1, self.nodes, self.origin)) - def writer(self, replacement: bool=False) -> 'Transaction': + def writer(self, replacement: bool = False) -> "Transaction": txn = Transaction(self, replacement) txn._setup_version() return txn - def origin_information(self) -> Tuple[Optional[dns.name.Name], bool, Optional[dns.name.Name]]: + def origin_information( + self, + ) -> Tuple[Optional[dns.name.Name], bool, Optional[dns.name.Name]]: effective: Optional[dns.name.Name] if self.relativize: effective = dns.name.empty @@ -839,8 +893,9 @@ class Zone(dns.transaction.TransactionManager): # A node with a version id. + class VersionedNode(dns.node.Node): # lgtm[py/missing-equals] - __slots__ = ['id'] + __slots__ = ["id"] def __init__(self): super().__init__() @@ -850,7 +905,7 @@ class VersionedNode(dns.node.Node): # lgtm[py/missing-equals] @dns.immutable.immutable class ImmutableVersionedNode(VersionedNode): - __slots__ = ['id'] + __slots__ = ["id"] def __init__(self, node): super().__init__() @@ -859,22 +914,34 @@ class ImmutableVersionedNode(VersionedNode): [dns.rdataset.ImmutableRdataset(rds) for rds in node.rdatasets] ) - def find_rdataset(self, rdclass: dns.rdataclass.RdataClass, rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> dns.rdataset.Rdataset: + def find_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> dns.rdataset.Rdataset: if create: raise TypeError("immutable") return super().find_rdataset(rdclass, rdtype, covers, False) - def get_rdataset(self, rdclass: dns.rdataclass.RdataClass, rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE, - create: bool=False) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + create: bool = False, + ) -> Optional[dns.rdataset.Rdataset]: if create: raise TypeError("immutable") return super().get_rdataset(rdclass, rdtype, covers, False) - def delete_rdataset(self, rdclass: dns.rdataclass.RdataClass, rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType=dns.rdatatype.NONE) -> None: + def delete_rdataset( + self, + rdclass: dns.rdataclass.RdataClass, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType = dns.rdatatype.NONE, + ) -> None: raise TypeError("immutable") def replace_rdataset(self, replacement: dns.rdataset.Rdataset) -> None: @@ -885,9 +952,13 @@ class ImmutableVersionedNode(VersionedNode): class Version: - def __init__(self, zone: Zone, id: int, - nodes: Optional[Dict[dns.name.Name, dns.node.Node]]=None, - origin: Optional[dns.name.Name]=None): + def __init__( + self, + zone: Zone, + id: int, + nodes: Optional[Dict[dns.name.Name, dns.node.Node]] = None, + origin: Optional[dns.name.Name] = None, + ): self.zone = zone self.id = id if nodes is not None: @@ -902,7 +973,7 @@ class Version: # This should probably never happen as other code (e.g. # _rr_line) will notice the lack of an origin before us, but # we check just in case! - raise KeyError('no zone origin is defined') + raise KeyError("no zone origin is defined") if not name.is_subdomain(self.origin): raise KeyError("name is not a subdomain of the zone origin") if self.zone.relativize: @@ -910,7 +981,7 @@ class Version: elif not self.zone.relativize: # We have a relative name in a non-relative zone, so derelativize. if self.origin is None: - raise KeyError('no zone origin is defined') + raise KeyError("no zone origin is defined") name = name.derelativize(self.origin) return name @@ -918,8 +989,12 @@ class Version: name = self._validate_name(name) return self.nodes.get(name) - def get_rdataset(self, name: dns.name.Name, rdtype: dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType) -> Optional[dns.rdataset.Rdataset]: + def get_rdataset( + self, + name: dns.name.Name, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType, + ) -> Optional[dns.rdataset.Rdataset]: node = self.get_node(name) if node is None: return None @@ -930,7 +1005,7 @@ class Version: class WritableVersion(Version): - def __init__(self, zone: Zone, replacement: bool=False): + def __init__(self, zone: Zone, replacement: bool = False): # The zone._versions_lock must be held by our caller in a versioned # zone. id = zone._get_next_version_id() @@ -951,14 +1026,14 @@ class WritableVersion(Version): node = self.nodes.get(name) if node is None or name not in self.changed: new_node = self.zone.node_factory() - if hasattr(new_node, 'id'): + if hasattr(new_node, "id"): # We keep doing this for backwards compatibility, as earlier # code used new_node.id != self.id for the "do we need to CoW?" # test. Now we use the changed set as this works with both # regular zones and versioned zones. # # We ignore the mypy error as this is safe but it doesn't see it. - new_node.id = self.id # type: ignore + new_node.id = self.id # type: ignore if node is not None: # moo! copy on write! new_node.rdatasets.extend(node.rdatasets) @@ -974,12 +1049,18 @@ class WritableVersion(Version): del self.nodes[name] self.changed.add(name) - def put_rdataset(self, name: dns.name.Name, rdataset: dns.rdataset.Rdataset) -> None: + def put_rdataset( + self, name: dns.name.Name, rdataset: dns.rdataset.Rdataset + ) -> None: node = self._maybe_cow(name) node.replace_rdataset(rdataset) - def delete_rdataset(self, name: dns.name.Name, rdtype:dns.rdatatype.RdataType, - covers: dns.rdatatype.RdataType) -> None: + def delete_rdataset( + self, + name: dns.name.Name, + rdtype: dns.rdatatype.RdataType, + covers: dns.rdatatype.RdataType, + ) -> None: node = self._maybe_cow(name) node.delete_rdataset(self.zone.rdclass, rdtype, covers) if len(node) == 0: @@ -1009,7 +1090,6 @@ class ImmutableVersion(Version): class Transaction(dns.transaction.Transaction): - def __init__(self, zone, replacement, version=None, make_immutable=False): read_only = version is not None super().__init__(zone, replacement, read_only) @@ -1086,11 +1166,17 @@ class Transaction(dns.transaction.Transaction): return (absolute, relativize, effective) -def from_text(text: str, origin: Optional[Union[dns.name.Name, str]]=None, - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, - relativize: bool=True, zone_factory: Any=Zone, filename: Optional[str]=None, - allow_include: bool=False, check_origin: bool=True, - idna_codec: Optional[dns.name.IDNACodec]=None) -> Zone: +def from_text( + text: str, + origin: Optional[Union[dns.name.Name, str]] = None, + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + relativize: bool = True, + zone_factory: Any = Zone, + filename: Optional[str] = None, + allow_include: bool = False, + check_origin: bool = True, + idna_codec: Optional[dns.name.IDNACodec] = None, +) -> Zone: """Build a zone object from a zone file format string. *text*, a ``str``, the zone file format input. @@ -1099,7 +1185,8 @@ def from_text(text: str, origin: Optional[Union[dns.name.Name, str]]=None, of the zone; if not specified, the first ``$ORIGIN`` statement in the zone file will determine the origin of the zone. - *rdclass*, a ``dns.rdataclass.RdataClass``, the zone's rdata class; the default is class IN. + *rdclass*, a ``dns.rdataclass.RdataClass``, the zone's rdata class; the default is + class IN. *relativize*, a ``bool``, determine's whether domain names are relativized to the zone's origin. The default is ``True``. @@ -1137,12 +1224,11 @@ def from_text(text: str, origin: Optional[Union[dns.name.Name, str]]=None, # interface is from_file(). if filename is None: - filename = '<string>' + filename = "<string>" zone = zone_factory(origin, rdclass, relativize=relativize) with zone.writer(True) as txn: tok = dns.tokenizer.Tokenizer(text, filename, idna_codec=idna_codec) - reader = dns.zonefile.Reader(tok, rdclass, txn, - allow_include=allow_include) + reader = dns.zonefile.Reader(tok, rdclass, txn, allow_include=allow_include) try: reader.read() except dns.zonefile.UnknownOrigin: @@ -1154,10 +1240,16 @@ def from_text(text: str, origin: Optional[Union[dns.name.Name, str]]=None, return zone -def from_file(f: Any, origin: Optional[Union[dns.name.Name, str]]=None, - rdclass: dns.rdataclass.RdataClass=dns.rdataclass.IN, - relativize: bool=True, zone_factory: Any=Zone, filename: Optional[str]=None, - allow_include: bool=True, check_origin: bool=True) -> Zone: +def from_file( + f: Any, + origin: Optional[Union[dns.name.Name, str]] = None, + rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN, + relativize: bool = True, + zone_factory: Any = Zone, + filename: Optional[str] = None, + allow_include: bool = True, + check_origin: bool = True, +) -> Zone: """Read a zone file and build a zone object. *f*, a file or ``str``. If *f* is a string, it is treated @@ -1205,12 +1297,25 @@ def from_file(f: Any, origin: Optional[Union[dns.name.Name, str]]=None, if filename is None: filename = f f = stack.enter_context(open(f)) - return from_text(f, origin, rdclass, relativize, zone_factory, - filename, allow_include, check_origin) + return from_text( + f, + origin, + rdclass, + relativize, + zone_factory, + filename, + allow_include, + check_origin, + ) assert False # make mypy happy lgtm[py/unreachable-statement] -def from_xfr(xfr: Any, zone_factory: Any=Zone, relativize: bool=True, check_origin: bool=True) -> Zone: +def from_xfr( + xfr: Any, + zone_factory: Any = Zone, + relativize: bool = True, + check_origin: bool = True, +) -> Zone: """Convert the output of a zone transfer generator into a zone object. *xfr*, a generator of ``dns.message.Message`` objects, typically @@ -1250,13 +1355,12 @@ def from_xfr(xfr: Any, zone_factory: Any=Zone, relativize: bool=True, check_orig if not znode: znode = z.node_factory() z.nodes[rrset.name] = znode - zrds = znode.find_rdataset(rrset.rdclass, rrset.rdtype, - rrset.covers, True) + zrds = znode.find_rdataset(rrset.rdclass, rrset.rdtype, rrset.covers, True) zrds.update_ttl(rrset.ttl) for rd in rrset: zrds.add(rd) if z is None: - raise ValueError('empty transfer') + raise ValueError("empty transfer") if check_origin: z.check_origin() return z diff --git a/dns/zonefile.py b/dns/zonefile.py index 479f0d6..fd17073 100644 --- a/dns/zonefile.py +++ b/dns/zonefile.py @@ -51,42 +51,53 @@ def _check_cname_and_other_data(txn, name, rdataset): # empty nodes are neutral. return node_kind = node.classify() - if node_kind == dns.node.NodeKind.CNAME and \ - rdataset_kind == dns.node.NodeKind.REGULAR: - raise CNAMEAndOtherData('rdataset type is not compatible with a ' - 'CNAME node') - elif node_kind == dns.node.NodeKind.REGULAR and \ - rdataset_kind == dns.node.NodeKind.CNAME: - raise CNAMEAndOtherData('CNAME rdataset is not compatible with a ' - 'regular data node') + if ( + node_kind == dns.node.NodeKind.CNAME + and rdataset_kind == dns.node.NodeKind.REGULAR + ): + raise CNAMEAndOtherData("rdataset type is not compatible with a " "CNAME node") + elif ( + node_kind == dns.node.NodeKind.REGULAR + and rdataset_kind == dns.node.NodeKind.CNAME + ): + raise CNAMEAndOtherData( + "CNAME rdataset is not compatible with a " "regular data node" + ) # Otherwise at least one of the node and the rdataset is neutral, so # adding the rdataset is ok -SavedStateType = Tuple[dns.tokenizer.Tokenizer, - Optional[dns.name.Name], # current_origin - Optional[dns.name.Name], # last_name - Optional[Any], # current_file - int, # last_ttl - bool, # last_ttl_known - int, # default_ttl - bool] # default_ttl_known +SavedStateType = Tuple[ + dns.tokenizer.Tokenizer, + Optional[dns.name.Name], # current_origin + Optional[dns.name.Name], # last_name + Optional[Any], # current_file + int, # last_ttl + bool, # last_ttl_known + int, # default_ttl + bool, +] # default_ttl_known class Reader: """Read a DNS zone file into a transaction.""" - def __init__(self, tok: dns.tokenizer.Tokenizer, rdclass: dns.rdataclass.RdataClass, - txn: dns.transaction.Transaction, allow_include: bool=False, - allow_directives: bool=True, force_name: Optional[dns.name.Name]=None, - force_ttl: Optional[int]=None, - force_rdclass: Optional[dns.rdataclass.RdataClass]=None, - force_rdtype: Optional[dns.rdatatype.RdataType]=None, - default_ttl: Optional[int]=None): + def __init__( + self, + tok: dns.tokenizer.Tokenizer, + rdclass: dns.rdataclass.RdataClass, + txn: dns.transaction.Transaction, + allow_include: bool = False, + allow_directives: bool = True, + force_name: Optional[dns.name.Name] = None, + force_ttl: Optional[int] = None, + force_rdclass: Optional[dns.rdataclass.RdataClass] = None, + force_rdtype: Optional[dns.rdatatype.RdataType] = None, + default_ttl: Optional[int] = None, + ): self.tok = tok - (self.zone_origin, self.relativize, _) = \ - txn.manager.origin_information() + (self.zone_origin, self.relativize, _) = txn.manager.origin_information() self.current_origin = self.zone_origin self.last_ttl = 0 self.last_ttl_known = False @@ -191,13 +202,17 @@ class Reader: try: rdtype = dns.rdatatype.from_text(token.value) except Exception: - raise dns.exception.SyntaxError( - "unknown rdatatype '%s'" % token.value) + raise dns.exception.SyntaxError("unknown rdatatype '%s'" % token.value) try: - rd = dns.rdata.from_text(rdclass, rdtype, self.tok, - self.current_origin, self.relativize, - self.zone_origin) + rd = dns.rdata.from_text( + rdclass, + rdtype, + self.tok, + self.current_origin, + self.relativize, + self.zone_origin, + ) except dns.exception.SyntaxError: # Catch and reraise. raise @@ -209,7 +224,8 @@ class Reader: # helpful filename:line info. (ty, va) = sys.exc_info()[:2] raise dns.exception.SyntaxError( - "caught exception {}: {}".format(str(ty), str(va))) + "caught exception {}: {}".format(str(ty), str(va)) + ) if not self.default_ttl_known and rdtype == dns.rdatatype.SOA: # The pre-RFC2308 and pre-BIND9 behavior inherits the zone default @@ -240,30 +256,30 @@ class Reader: g1 = is_generate1.match(side) if g1: mod, sign, offset, width, base = g1.groups() - if sign == '': - sign = '+' + if sign == "": + sign = "+" g2 = is_generate2.match(side) if g2: mod, sign, offset = g2.groups() - if sign == '': - sign = '+' + if sign == "": + sign = "+" width = 0 - base = 'd' + base = "d" g3 = is_generate3.match(side) if g3: mod, sign, offset, width = g3.groups() - if sign == '': - sign = '+' - base = 'd' + if sign == "": + sign = "+" + base = "d" if not (g1 or g2 or g3): - mod = '' - sign = '+' + mod = "" + sign = "+" offset = 0 width = 0 - base = 'd' + base = "d" - if base != 'd': + if base != "d": raise NotImplementedError() return mod, sign, offset, width, base @@ -328,8 +344,7 @@ class Reader: if not token.is_identifier(): raise dns.exception.SyntaxError except Exception: - raise dns.exception.SyntaxError("unknown rdatatype '%s'" % - token.value) + raise dns.exception.SyntaxError("unknown rdatatype '%s'" % token.value) # rhs (required) rhs = token.value @@ -341,24 +356,25 @@ class Reader: for i in range(start, stop + 1, step): # +1 because bind is inclusive and python is exclusive - if lsign == '+': + if lsign == "+": lindex = i + int(loffset) - elif lsign == '-': + elif lsign == "-": lindex = i - int(loffset) - if rsign == '-': + if rsign == "-": rindex = i - int(roffset) - elif rsign == '+': + elif rsign == "+": rindex = i + int(roffset) lzfindex = str(lindex).zfill(int(lwidth)) rzfindex = str(rindex).zfill(int(rwidth)) - name = lhs.replace('$%s' % (lmod), lzfindex) - rdata = rhs.replace('$%s' % (rmod), rzfindex) + name = lhs.replace("$%s" % (lmod), lzfindex) + rdata = rhs.replace("$%s" % (rmod), rzfindex) - self.last_name = dns.name.from_text(name, self.current_origin, - self.tok.idna_codec) + self.last_name = dns.name.from_text( + name, self.current_origin, self.tok.idna_codec + ) name = self.last_name if not name.is_subdomain(self.zone_origin): self._eat_line() @@ -367,9 +383,14 @@ class Reader: name = name.relativize(self.zone_origin) try: - rd = dns.rdata.from_text(rdclass, rdtype, rdata, - self.current_origin, self.relativize, - self.zone_origin) + rd = dns.rdata.from_text( + rdclass, + rdtype, + rdata, + self.current_origin, + self.relativize, + self.zone_origin, + ) except dns.exception.SyntaxError: # Catch and reraise. raise @@ -380,8 +401,9 @@ class Reader: # We convert them to syntax errors so that we can emit # helpful filename:line info. (ty, va) = sys.exc_info()[:2] - raise dns.exception.SyntaxError("caught exception %s: %s" % - (str(ty), str(va))) + raise dns.exception.SyntaxError( + "caught exception %s: %s" % (str(ty), str(va)) + ) self.txn.add(name, ttl, rd) @@ -399,14 +421,16 @@ class Reader: if self.current_file is not None: self.current_file.close() if len(self.saved_state) > 0: - (self.tok, - self.current_origin, - self.last_name, - self.current_file, - self.last_ttl, - self.last_ttl_known, - self.default_ttl, - self.default_ttl_known) = self.saved_state.pop(-1) + ( + self.tok, + self.current_origin, + self.last_name, + self.current_file, + self.last_ttl, + self.last_ttl_known, + self.default_ttl, + self.default_ttl_known, + ) = self.saved_state.pop(-1) continue break elif token.is_eol(): @@ -414,51 +438,56 @@ class Reader: elif token.is_comment(): self.tok.get_eol() continue - elif token.value[0] == '$' and self.allow_directives: + elif token.value[0] == "$" and self.allow_directives: c = token.value.upper() - if c == '$TTL': + if c == "$TTL": token = self.tok.get() if not token.is_identifier(): raise dns.exception.SyntaxError("bad $TTL") self.default_ttl = dns.ttl.from_text(token.value) self.default_ttl_known = True self.tok.get_eol() - elif c == '$ORIGIN': + elif c == "$ORIGIN": self.current_origin = self.tok.get_name() self.tok.get_eol() if self.zone_origin is None: self.zone_origin = self.current_origin self.txn._set_origin(self.current_origin) - elif c == '$INCLUDE' and self.allow_include: + elif c == "$INCLUDE" and self.allow_include: token = self.tok.get() filename = token.value token = self.tok.get() new_origin: Optional[dns.name.Name] if token.is_identifier(): - new_origin = dns.name.from_text(token.value, self.current_origin, self.tok.idna_codec) + new_origin = dns.name.from_text( + token.value, self.current_origin, self.tok.idna_codec + ) self.tok.get_eol() elif not token.is_eol_or_eof(): - raise dns.exception.SyntaxError( - "bad origin in $INCLUDE") + raise dns.exception.SyntaxError("bad origin in $INCLUDE") else: new_origin = self.current_origin - self.saved_state.append((self.tok, - self.current_origin, - self.last_name, - self.current_file, - self.last_ttl, - self.last_ttl_known, - self.default_ttl, - self.default_ttl_known)) - self.current_file = open(filename, 'r') - self.tok = dns.tokenizer.Tokenizer(self.current_file, - filename) + self.saved_state.append( + ( + self.tok, + self.current_origin, + self.last_name, + self.current_file, + self.last_ttl, + self.last_ttl_known, + self.default_ttl, + self.default_ttl_known, + ) + ) + self.current_file = open(filename, "r") + self.tok = dns.tokenizer.Tokenizer(self.current_file, filename) self.current_origin = new_origin - elif c == '$GENERATE': + elif c == "$GENERATE": self._generate_line() else: raise dns.exception.SyntaxError( - "Unknown zone file directive '" + c + "'") + "Unknown zone file directive '" + c + "'" + ) continue self.tok.unget(token) self._rr_line() @@ -467,13 +496,13 @@ class Reader: if detail is None: detail = "syntax error" ex = dns.exception.SyntaxError( - "%s:%d: %s" % (filename, line_number, detail)) + "%s:%d: %s" % (filename, line_number, detail) + ) tb = sys.exc_info()[2] raise ex.with_traceback(tb) from None class RRsetsReaderTransaction(dns.transaction.Transaction): - def __init__(self, manager, replacement, read_only): assert not read_only super().__init__(manager, replacement, read_only) @@ -525,8 +554,9 @@ class RRsetsReaderTransaction(dns.transaction.Transaction): if commit and self._changed(): rrsets = [] for (name, _, _), rdataset in self.rdatasets.items(): - rrset = dns.rrset.RRset(name, rdataset.rdclass, rdataset.rdtype, - rdataset.covers) + rrset = dns.rrset.RRset( + name, rdataset.rdclass, rdataset.rdtype, rdataset.covers + ) rrset.update(rdataset) rrsets.append(rrset) self.manager.set_rrsets(rrsets) @@ -536,8 +566,9 @@ class RRsetsReaderTransaction(dns.transaction.Transaction): class RRSetsReaderManager(dns.transaction.TransactionManager): - def __init__(self, origin=dns.name.root, relativize=False, - rdclass=dns.rdataclass.IN): + def __init__( + self, origin=dns.name.root, relativize=False, rdclass=dns.rdataclass.IN + ): self.origin = origin self.relativize = relativize self.rdclass = rdclass @@ -561,16 +592,18 @@ class RRSetsReaderManager(dns.transaction.TransactionManager): self.rrsets = rrsets -def read_rrsets(text: Any, - name: Optional[Union[dns.name.Name, str]]=None, - ttl: Optional[int]=None, - rdclass: Optional[Union[dns.rdataclass.RdataClass, str]]=dns.rdataclass.IN, - default_rdclass: Union[dns.rdataclass.RdataClass, str]=dns.rdataclass.IN, - rdtype: Optional[Union[dns.rdatatype.RdataType, str]]=None, - default_ttl: Optional[Union[int, str]]=None, - idna_codec: Optional[dns.name.IDNACodec]=None, - origin: Optional[Union[dns.name.Name, str]]=dns.name.root, - relativize: bool=False) -> List[dns.rrset.RRset]: +def read_rrsets( + text: Any, + name: Optional[Union[dns.name.Name, str]] = None, + ttl: Optional[int] = None, + rdclass: Optional[Union[dns.rdataclass.RdataClass, str]] = dns.rdataclass.IN, + default_rdclass: Union[dns.rdataclass.RdataClass, str] = dns.rdataclass.IN, + rdtype: Optional[Union[dns.rdatatype.RdataType, str]] = None, + default_ttl: Optional[Union[int, str]] = None, + idna_codec: Optional[dns.name.IDNACodec] = None, + origin: Optional[Union[dns.name.Name, str]] = dns.name.root, + relativize: bool = False, +) -> List[dns.rrset.RRset]: """Read one or more rrsets from the specified text, possibly subject to restrictions. @@ -639,9 +672,17 @@ def read_rrsets(text: Any, the_rdtype = None manager = RRSetsReaderManager(origin, relativize, default_rdclass) with manager.writer(True) as txn: - tok = dns.tokenizer.Tokenizer(text, '<input>', idna_codec=idna_codec) - reader = Reader(tok, the_default_rdclass, txn, allow_directives=False, - force_name=name, force_ttl=ttl, force_rdclass=the_rdclass, - force_rdtype=the_rdtype, default_ttl=default_ttl) + tok = dns.tokenizer.Tokenizer(text, "<input>", idna_codec=idna_codec) + reader = Reader( + tok, + the_default_rdclass, + txn, + allow_directives=False, + force_name=name, + force_ttl=ttl, + force_rdclass=the_rdclass, + force_rdtype=the_rdtype, + default_ttl=default_ttl, + ) reader.read() return manager.rrsets |
