summaryrefslogtreecommitdiff
path: root/dns
diff options
context:
space:
mode:
authorBob Halley <halley@dnspython.org>2022-03-15 08:37:20 -0700
committerBob Halley <halley@dnspython.org>2022-03-15 08:37:20 -0700
commitb1d2332687adbecc0acbb4e623124f783f859d9e (patch)
tree5318d5ecc0dd35e0a6922380cd60f9d9caa9ad34 /dns
parent08f8bde64e8679d5e4f0b129292461de152ba32b (diff)
downloaddnspython-b1d2332687adbecc0acbb4e623124f783f859d9e.tar.gz
black autoformatting
Diffstat (limited to 'dns')
-rw-r--r--dns/__init__.py90
-rw-r--r--dns/_asyncbackend.py22
-rw-r--r--dns/_asyncio_backend.py75
-rw-r--r--dns/_curio_backend.py38
-rw-r--r--dns/_immutable_ctx.py8
-rw-r--r--dns/_trio_backend.py32
-rw-r--r--dns/asyncbackend.py33
-rw-r--r--dns/asyncquery.py400
-rw-r--r--dns/asyncresolver.py136
-rw-r--r--dns/dnssec.py171
-rw-r--r--dns/e164.py30
-rw-r--r--dns/edns.py101
-rw-r--r--dns/entropy.py16
-rw-r--r--dns/enum.py7
-rw-r--r--dns/exception.py20
-rw-r--r--dns/flags.py5
-rw-r--r--dns/grange.py15
-rw-r--r--dns/immutable.py4
-rw-r--r--dns/inet.py10
-rw-r--r--dns/ipv4.py11
-rw-r--r--dns/ipv6.py68
-rw-r--r--dns/message.py653
-rw-r--r--dns/name.py228
-rw-r--r--dns/namedict.py2
-rw-r--r--dns/node.py117
-rw-r--r--dns/opcode.py2
-rw-r--r--dns/query.py539
-rw-r--r--dns/rcode.py14
-rw-r--r--dns/rdata.py340
-rw-r--r--dns/rdataclass.py3
-rw-r--r--dns/rdataset.py192
-rw-r--r--dns/rdatatype.py22
-rw-r--r--dns/rdtypes/ANY/AMTRELAY.py48
-rw-r--r--dns/rdtypes/ANY/CAA.py19
-rw-r--r--dns/rdtypes/ANY/CDNSKEY.py8
-rw-r--r--dns/rdtypes/ANY/CERT.py72
-rw-r--r--dns/rdtypes/ANY/CSYNC.py13
-rw-r--r--dns/rdtypes/ANY/DNSKEY.py8
-rw-r--r--dns/rdtypes/ANY/GPOS.py42
-rw-r--r--dns/rdtypes/ANY/HINFO.py16
-rw-r--r--dns/rdtypes/ANY/HIP.py17
-rw-r--r--dns/rdtypes/ANY/ISDN.py20
-rw-r--r--dns/rdtypes/ANY/L32.py11
-rw-r--r--dns/rdtypes/ANY/L64.py20
-rw-r--r--dns/rdtypes/ANY/LOC.py176
-rw-r--r--dns/rdtypes/ANY/LP.py11
-rw-r--r--dns/rdtypes/ANY/NID.py19
-rw-r--r--dns/rdtypes/ANY/NSEC.py11
-rw-r--r--dns/rdtypes/ANY/NSEC3.py61
-rw-r--r--dns/rdtypes/ANY/NSEC3PARAM.py23
-rw-r--r--dns/rdtypes/ANY/OPENPGPKEY.py6
-rw-r--r--dns/rdtypes/ANY/OPT.py11
-rw-r--r--dns/rdtypes/ANY/RP.py7
-rw-r--r--dns/rdtypes/ANY/RRSIG.py79
-rw-r--r--dns/rdtypes/ANY/SOA.py37
-rw-r--r--dns/rdtypes/ANY/SSHFP.py22
-rw-r--r--dns/rdtypes/ANY/TKEY.py60
-rw-r--r--dns/rdtypes/ANY/TSIG.py91
-rw-r--r--dns/rdtypes/ANY/URI.py14
-rw-r--r--dns/rdtypes/ANY/X25.py9
-rw-r--r--dns/rdtypes/ANY/ZONEMD.py28
-rw-r--r--dns/rdtypes/ANY/__init__.py94
-rw-r--r--dns/rdtypes/CH/A.py10
-rw-r--r--dns/rdtypes/CH/__init__.py2
-rw-r--r--dns/rdtypes/IN/A.py7
-rw-r--r--dns/rdtypes/IN/AAAA.py7
-rw-r--r--dns/rdtypes/IN/APL.py32
-rw-r--r--dns/rdtypes/IN/DHCID.py7
-rw-r--r--dns/rdtypes/IN/HTTPS.py1
-rw-r--r--dns/rdtypes/IN/IPSECKEY.py53
-rw-r--r--dns/rdtypes/IN/NAPTR.py48
-rw-r--r--dns/rdtypes/IN/NSAP.py15
-rw-r--r--dns/rdtypes/IN/PX.py9
-rw-r--r--dns/rdtypes/IN/SRV.py12
-rw-r--r--dns/rdtypes/IN/SVCB.py1
-rw-r--r--dns/rdtypes/IN/WKS.py18
-rw-r--r--dns/rdtypes/IN/__init__.py28
-rw-r--r--dns/rdtypes/__init__.py24
-rw-r--r--dns/rdtypes/dnskeybase.py24
-rw-r--r--dns/rdtypes/dsbase.py36
-rw-r--r--dns/rdtypes/euibase.py26
-rw-r--r--dns/rdtypes/mxbase.py9
-rw-r--r--dns/rdtypes/nsbase.py7
-rw-r--r--dns/rdtypes/svcbbase.py129
-rw-r--r--dns/rdtypes/tlsabase.py24
-rw-r--r--dns/rdtypes/txtbase.py52
-rw-r--r--dns/rdtypes/util.py34
-rw-r--r--dns/renderer.py78
-rw-r--r--dns/resolver.py642
-rw-r--r--dns/reversename.py33
-rw-r--r--dns/rrset.py134
-rw-r--r--dns/serial.py31
-rw-r--r--dns/set.py37
-rw-r--r--dns/tokenizer.py171
-rw-r--r--dns/transaction.py138
-rw-r--r--dns/tsig.py62
-rw-r--r--dns/tsigkeyring.py2
-rw-r--r--dns/ttl.py12
-rw-r--r--dns/update.py156
-rw-r--r--dns/version.py31
-rw-r--r--dns/versioned.py103
-rw-r--r--dns/win32util.py71
-rw-r--r--dns/wire.py19
-rw-r--r--dns/xfr.py113
-rw-r--r--dns/zone.py392
-rw-r--r--dns/zonefile.py253
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
diff --git a/dns/set.py b/dns/set.py
index a4e12b6..fa50ed9 100644
--- a/dns/set.py
+++ b/dns/set.py
@@ -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):
diff --git a/dns/ttl.py b/dns/ttl.py
index 9f5730e..264b033 100644
--- a/dns/ttl.py
+++ b/dns/ttl.py
@@ -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)
diff --git a/dns/xfr.py b/dns/xfr.py
index a360deb..89e92ca 100644
--- a/dns/xfr.py
+++ b/dns/xfr.py
@@ -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