Загрузить файлы в «venv/Lib/site-packages/dns»

This commit is contained in:
2026-07-02 18:08:01 +00:00
parent 0bb539233f
commit fa17e96ca7
5 changed files with 1279 additions and 0 deletions

View File

@@ -0,0 +1,90 @@
# Copyright (C) Dnspython Contributors, see LICENSE for text of ISC license
# Copyright (C) 2003-2017 Nominum, Inc.
#
# Permission to use, copy, modify, and distribute this software and its
# documentation for any purpose with or without fee is hereby granted,
# provided that the above copyright notice and this permission notice
# appear in all copies.
#
# THE SOFTWARE IS PROVIDED "AS IS" AND NOMINUM DISCLAIMS ALL WARRANTIES
# WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
# MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL NOMINUM BE LIABLE FOR
# ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
# WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
# ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT
# OF OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
"""DNS TTL conversion."""
import dns.exception
# Technically TTLs are supposed to be between 0 and 2**31 - 1, with values
# greater than that interpreted as 0, but we do not impose this policy here
# as values > 2**31 - 1 occur in real world data.
#
# We leave it to applications to impose tighter bounds if desired.
MAX_TTL = 2**32 - 1
class BadTTL(dns.exception.SyntaxError):
"""DNS TTL value is not well-formed."""
def from_text(text: str) -> int:
"""Convert the text form of a TTL to an integer.
The BIND 8 units syntax for TTLs (e.g. '1w6d4h3m10s') is supported.
*text*, a ``str``, the textual TTL.
Raises ``dns.ttl.BadTTL`` if the TTL is not well-formed.
Returns an ``int``.
"""
if text.isdigit():
total = int(text)
elif len(text) == 0:
raise BadTTL
else:
total = 0
current = 0
need_digit = True
for c in text:
if c.isdigit():
current *= 10
current += int(c)
need_digit = False
else:
if need_digit:
raise BadTTL
c = c.lower()
if c == "w":
total += current * 604800
elif c == "d":
total += current * 86400
elif c == "h":
total += current * 3600
elif c == "m":
total += current * 60
elif c == "s":
total += current
else:
raise BadTTL(f"unknown unit '{c}'")
current = 0
need_digit = True
if not current == 0:
raise BadTTL("trailing integer")
if total < 0 or total > MAX_TTL:
raise BadTTL("TTL should be between 0 and 2**32 - 1 (inclusive)")
return total
def make(value: int | str) -> int:
if isinstance(value, int):
return value
elif isinstance(value, str):
return from_text(value)
else:
raise ValueError("cannot convert value to TTL")

View File

@@ -0,0 +1,389 @@
# Copyright (C) Dnspython Contributors, see LICENSE for text of ISC license
# Copyright (C) 2003-2007, 2009-2011 Nominum, Inc.
#
# Permission to use, copy, modify, and distribute this software and its
# documentation for any purpose with or without fee is hereby granted,
# provided that the above copyright notice and this permission notice
# appear in all copies.
#
# THE SOFTWARE IS PROVIDED "AS IS" AND NOMINUM DISCLAIMS ALL WARRANTIES
# WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
# MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL NOMINUM BE LIABLE FOR
# ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
# WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
# ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT
# OF OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
"""DNS Dynamic Update Support"""
from typing import Any, List
import dns.enum
import dns.exception
import dns.message
import dns.name
import dns.opcode
import dns.rdata
import dns.rdataclass
import dns.rdataset
import dns.rdatatype
import dns.rrset
import dns.tsig
class UpdateSection(dns.enum.IntEnum):
"""Update sections"""
ZONE = 0
PREREQ = 1
UPDATE = 2
ADDITIONAL = 3
@classmethod
def _maximum(cls):
return 3
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: dns.name.Name | str | None = None,
rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN,
keyring: Any | None = None,
keyname: dns.name.Name | None = None,
keyalgorithm: dns.name.Name | str = dns.tsig.default_algorithm,
id: int | None = None,
):
"""Initialize a new DNS Update object.
See the documentation of the Message class for a complete
description of the keyring dictionary.
*zone*, a ``dns.name.Name``, ``str``, or ``None``, the zone
which is being updated. ``None`` should only be used by dnspython's
message constructors, as a zone is required for the convenience
methods like ``add()``, ``replace()``, etc.
*rdclass*, an ``int`` or ``str``, the class of the zone.
The *keyring*, *keyname*, and *keyalgorithm* parameters are passed to
``use_tsig()``; see its documentation for details.
"""
super().__init__(id=id)
self.flags |= dns.opcode.to_flags(dns.opcode.UPDATE)
if isinstance(zone, str):
zone = dns.name.from_text(zone)
self.origin = zone
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,
)
if keyring is not None:
self.use_tsig(keyring, keyname, algorithm=keyalgorithm)
@property
def zone(self) -> List[dns.rrset.RRset]:
"""The zone section."""
return self.sections[0]
@zone.setter
def zone(self, v):
self.sections[0] = v
@property
def prerequisite(self) -> List[dns.rrset.RRset]:
"""The prerequisite section."""
return self.sections[1]
@prerequisite.setter
def prerequisite(self, v):
self.sections[1] = v
@property
def update(self) -> List[dns.rrset.RRset]:
"""The update section."""
return self.sections[2]
@update.setter
def update(self, v):
self.sections[2] = v
def _add_rr(self, name, ttl, rd, deleting=None, section=None):
"""Add a single RR to the update section."""
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.add(rd, ttl)
def _add(self, replace, section, name, *args):
"""Add records.
*replace* is the replacement mode. If ``False``,
RRs are added to an existing RRset; if ``True``, the RRset
is replaced with the specified contents. The second
argument is the section to add to. The third argument
is always a name. The other arguments can be:
- rdataset...
- ttl, rdata...
- ttl, rdtype, string...
"""
if isinstance(name, str):
name = dns.name.from_text(name, None)
if isinstance(args[0], dns.rdataset.Rdataset):
for rds in args:
if replace:
self.delete(name, rds.rdtype)
for rd in rds:
self._add_rr(name, rds.ttl, rd, section=section)
else:
args = list(args)
ttl = int(args.pop(0))
if isinstance(args[0], dns.rdata.Rdata):
if replace:
self.delete(name, args[0].rdtype)
for rd in args:
self._add_rr(name, ttl, rd, section=section)
else:
rdtype = dns.rdatatype.RdataType.make(args.pop(0))
if replace:
self.delete(name, rdtype)
for s in args:
rd = dns.rdata.from_text(self.zone_rdclass, rdtype, s, self.origin)
self._add_rr(name, ttl, rd, section=section)
def add(self, name: dns.name.Name | str, *args: Any) -> None:
"""Add records.
The first argument is always a name. The other
arguments can be:
- rdataset...
- ttl, rdata...
- ttl, rdtype, string...
"""
self._add(False, self.update, name, *args)
def delete(self, name: dns.name.Name | str, *args: Any) -> None:
"""Delete records.
The first argument is always a name. The other
arguments can be:
- *empty*
- rdataset...
- rdata...
- rdtype, [string...]
"""
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,
)
elif isinstance(args[0], dns.rdataset.Rdataset):
for rds in args:
for rd in rds:
self._add_rr(name, 0, rd, dns.rdataclass.NONE)
else:
largs = list(args)
if isinstance(largs[0], dns.rdata.Rdata):
for rd in largs:
self._add_rr(name, 0, rd, dns.rdataclass.NONE)
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,
)
else:
for s in largs:
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: dns.name.Name | str, *args: Any) -> None:
"""Replace records.
The first argument is always a name. The other
arguments can be:
- rdataset...
- ttl, rdata...
- ttl, rdtype, string...
Note that if you want to replace the entire node, you should do
a delete of the name followed by one or more calls to add.
"""
self._add(True, self.update, name, *args)
def present(self, name: dns.name.Name | str, *args: Any) -> None:
"""Require that an owner name (and optionally an rdata type,
or specific rdataset) exists as a prerequisite to the
execution of the update.
The first argument is always a name.
The other arguments can be:
- rdataset...
- rdata...
- rdtype, string...
"""
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
):
if not isinstance(args[0], dns.rdataset.Rdataset):
# Add a 0 TTL
largs = list(args)
largs.insert(0, 0) # type: ignore[arg-type]
self._add(False, self.prerequisite, name, *largs)
else:
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: dns.name.Name | str,
rdtype: dns.rdatatype.RdataType | str | None = 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,
)
else:
rdtype = dns.rdatatype.RdataType.make(rdtype)
self.find_rrset(
self.prerequisite,
name,
dns.rdataclass.NONE,
rdtype,
dns.rdatatype.NONE,
None,
True,
True,
)
def _get_one_rr_per_rrset(self, value):
# Updates are always one_rr_per_rrset
return True
def _parse_rr_header(self, section, name, rdclass, rdtype): # pyright: ignore
deleting = None
empty = False
if section == UpdateSection.ZONE:
if (
dns.rdataclass.is_metaclass(rdclass)
or rdtype != dns.rdatatype.SOA
or self.zone
):
raise dns.exception.FormError
else:
if not self.zone:
raise dns.exception.FormError
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
)
return (rdclass, rdtype, deleting, empty)
# backwards compatibility
Update = UpdateMessage
### BEGIN generated UpdateSection constants
ZONE = UpdateSection.ZONE
PREREQ = UpdateSection.PREREQ
UPDATE = UpdateSection.UPDATE
ADDITIONAL = UpdateSection.ADDITIONAL
### END generated UpdateSection constants

View File

@@ -0,0 +1,42 @@
# Copyright (C) Dnspython Contributors, see LICENSE for text of ISC license
# Copyright (C) 2003-2017 Nominum, Inc.
#
# Permission to use, copy, modify, and distribute this software and its
# documentation for any purpose with or without fee is hereby granted,
# provided that the above copyright notice and this permission notice
# appear in all copies.
#
# THE SOFTWARE IS PROVIDED "AS IS" AND NOMINUM DISCLAIMS ALL WARRANTIES
# WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
# MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL NOMINUM BE LIABLE FOR
# ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
# WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
# ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT
# OF OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
"""dnspython release version information."""
#: MAJOR
MAJOR = 2
#: MINOR
MINOR = 8
#: MICRO
MICRO = 0
#: RELEASELEVEL
RELEASELEVEL = 0x0F
#: SERIAL
SERIAL = 0
if RELEASELEVEL == 0x0F: # pragma: no cover lgtm[py/unreachable-statement]
#: version
version = f"{MAJOR}.{MINOR}.{MICRO}" # lgtm[py/unreachable-statement]
elif RELEASELEVEL == 0x00: # pragma: no cover lgtm[py/unreachable-statement]
version = f"{MAJOR}.{MINOR}.{MICRO}dev{SERIAL}" # lgtm[py/unreachable-statement]
elif RELEASELEVEL == 0x0C: # pragma: no cover lgtm[py/unreachable-statement]
version = f"{MAJOR}.{MINOR}.{MICRO}rc{SERIAL}" # lgtm[py/unreachable-statement]
else: # pragma: no cover lgtm[py/unreachable-statement]
version = f"{MAJOR}.{MINOR}.{MICRO}{RELEASELEVEL:x}{SERIAL}" # lgtm[py/unreachable-statement]
#: hexversion
hexversion = MAJOR << 24 | MINOR << 16 | MICRO << 8 | RELEASELEVEL << 4 | SERIAL

View File

@@ -0,0 +1,320 @@
# Copyright (C) Dnspython Contributors, see LICENSE for text of ISC license
"""DNS Versioned Zones."""
import collections
import threading
from typing import Callable, Deque, Set, cast
import dns.exception
import dns.name
import dns.node
import dns.rdataclass
import dns.rdataset
import dns.rdatatype
import dns.rdtypes.ANY.SOA
import dns.zone
class UseTransaction(dns.exception.DNSException):
"""To alter a versioned zone, use a transaction."""
# Backwards compatibility
Node = dns.zone.VersionedNode
ImmutableNode = dns.zone.ImmutableVersionedNode
Version = dns.zone.Version
WritableVersion = dns.zone.WritableVersion
ImmutableVersion = dns.zone.ImmutableVersion
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",
]
node_factory: Callable[[], dns.node.Node] = Node
def __init__(
self,
origin: dns.name.Name | str | None,
rdclass: dns.rdataclass.RdataClass = dns.rdataclass.IN,
relativize: bool = True,
pruning_policy: Callable[["Zone", Version], bool | None] | None = None,
):
"""Initialize a versioned zone object.
*origin* is the origin of the zone. It may be a ``dns.name.Name``,
a ``str``, or ``None``. If ``None``, then the zone's origin will
be set by the first ``$ORIGIN`` line in a zone file.
*rdclass*, an ``int``, 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``.
*pruning policy*, a function taking a ``Zone`` and a ``Version`` and returning
a ``bool``, or ``None``. Should the version be pruned? If ``None``,
the default policy, which retains one version is used.
"""
super().__init__(origin, rdclass, relativize)
self._versions: Deque[Version] = collections.deque()
self._version_lock = threading.Lock()
if pruning_policy is None:
self._pruning_policy = self._default_pruning_policy
else:
self._pruning_policy = pruning_policy
self._write_txn: Transaction | None = None
self._write_event: threading.Event | None = None
self._write_waiters: Deque[threading.Event] = collections.deque()
self._readers: Set[Transaction] = set()
self._commit_version_unlocked(
None, WritableVersion(self, replacement=True), origin
)
def reader(
self, id: int | None = None, serial: int | None = None
) -> Transaction: # pylint: disable=arguments-differ
if id is not None and serial is not None:
raise ValueError("cannot specify both id and serial")
with self._version_lock:
if id is not None:
version = None
for v in reversed(self._versions):
if v.id == id:
version = v
break
if version is None:
raise KeyError("version not found")
elif serial is not None:
if self.relativize:
oname = dns.name.empty
else:
assert self.origin is not None
oname = self.origin
version = None
for v in reversed(self._versions):
n = v.nodes.get(oname)
if n:
rds = n.get_rdataset(self.rdclass, dns.rdatatype.SOA)
if rds is None:
continue
soa = cast(dns.rdtypes.ANY.SOA.SOA, rds[0])
if rds and soa.serial == serial:
version = v
break
if version is None:
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:
event = None
while True:
with self._version_lock:
# Checking event == self._write_event ensures that either
# no one was waiting before we got lucky and found no write
# txn, or we were the one who was waiting and got woken up.
# This prevents "taking cuts" when creating a write txn.
if self._write_txn is None and event == self._write_event:
# Creating the transaction defers version setup
# (i.e. copying the nodes dictionary) until we
# 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
)
# give up our exclusive right to make a Transaction
self._write_event = None
break
# Someone else is writing already, so we will have to
# wait, but we want to do the actual wait outside the
# lock.
event = threading.Event()
self._write_waiters.append(event)
# wait (note we gave up the lock!)
#
# We only wake one sleeper at a time, so it's important
# that no event waiter can exit this method (e.g. via
# cancellation) without returning a transaction or waking
# someone else up.
#
# This is not a problem with Threading module threads as
# they cannot be canceled, but could be an issue with trio
# tasks when we do the async version of writer().
# I.e. we'd need to do something like:
#
# try:
# event.wait()
# except trio.Cancelled:
# with self._version_lock:
# self._maybe_wakeup_one_waiter_unlocked()
# raise
#
event.wait()
# Do the deferred version setup.
self._write_txn._setup_version()
return self._write_txn
def _maybe_wakeup_one_waiter_unlocked(self):
if len(self._write_waiters) > 0:
self._write_event = self._write_waiters.popleft()
self._write_event.set()
# pylint: disable=unused-argument
def _default_pruning_policy(self, zone, version):
return True
# pylint: enable=unused-argument
def _prune_versions_unlocked(self):
assert len(self._versions) > 0
# Don't ever prune a version greater than or equal to one that
# a reader has open. This pins versions in memory while the
# reader is open, and importantly lets the reader open a txn on
# a successor version (e.g. if generating an IXFR).
#
# Note our definition of least_kept also ensures we do not try to
# delete the greatest version.
if len(self._readers) > 0:
least_kept = min(txn.version.id for txn in self._readers) # pyright: ignore
else:
least_kept = self._versions[-1].id
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: int | None) -> None:
"""Set a pruning policy that retains up to the specified number
of versions
"""
if max_versions is not None and max_versions < 1:
raise ValueError("max versions must be at least 1")
if max_versions is None:
# pylint: disable=unused-argument
def policy(zone, _): # pyright: ignore
return False
else:
def policy(zone, _):
return len(zone._versions) > max_versions
self.set_pruning_policy(policy)
def set_pruning_policy(
self, policy: Callable[["Zone", Version], bool | None] | None
) -> None:
"""Set the pruning policy for the zone.
The *policy* function takes a `Version` and returns `True` if
the version should be pruned, and `False` otherwise. `None`
may also be specified for policy, in which case the default policy
is used.
Pruning checking proceeds from the least version and the first
time the function returns `False`, the checking stops. I.e. the
retained versions are always a consecutive sequence.
"""
if policy is None:
policy = self._default_pruning_policy
with self._version_lock:
self._pruning_policy = policy
self._prune_versions_unlocked()
def _end_read(self, txn):
with self._version_lock:
self._readers.remove(txn)
self._prune_versions_unlocked()
def _end_write_unlocked(self, txn):
assert self._write_txn == txn
self._write_txn = None
self._maybe_wakeup_one_waiter_unlocked()
def _end_write(self, txn):
with self._version_lock:
self._end_write_unlocked(txn)
def _commit_version_unlocked(self, txn, version, origin):
self._versions.append(version)
self._prune_versions_unlocked()
self.nodes = version.nodes
if self.origin is None:
self.origin = origin
# txn can be None in __init__ when we make the empty version.
if txn is not None:
self._end_write_unlocked(txn)
def _commit_version(self, txn, version, origin):
with self._version_lock:
self._commit_version_unlocked(txn, version, origin)
def _get_next_version_id(self):
if len(self._versions) > 0:
id = self._versions[-1].id + 1
else:
id = 1
return id
def find_node(
self, name: dns.name.Name | str, create: bool = False
) -> dns.node.Node:
if create:
raise UseTransaction
return super().find_node(name)
def delete_node(self, name: dns.name.Name | str) -> None:
raise UseTransaction
def find_rdataset(
self,
name: dns.name.Name | str,
rdtype: dns.rdatatype.RdataType | str,
covers: 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: dns.name.Name | str,
rdtype: dns.rdatatype.RdataType | str,
covers: dns.rdatatype.RdataType | str = dns.rdatatype.NONE,
create: bool = False,
) -> dns.rdataset.Rdataset | None:
if create:
raise UseTransaction
rdataset = super().get_rdataset(name, rdtype, covers)
if rdataset is not None:
return dns.rdataset.ImmutableRdataset(rdataset)
else:
return None
def delete_rdataset(
self,
name: dns.name.Name | str,
rdtype: dns.rdatatype.RdataType | str,
covers: dns.rdatatype.RdataType | str = dns.rdatatype.NONE,
) -> None:
raise UseTransaction
def replace_rdataset(
self, name: dns.name.Name | str, replacement: dns.rdataset.Rdataset
) -> None:
raise UseTransaction

View File

@@ -0,0 +1,438 @@
import sys
import dns._features
# pylint: disable=W0612,W0613,C0301
if sys.platform == "win32":
import ctypes
import ctypes.wintypes as wintypes
import winreg # pylint: disable=import-error
from enum import IntEnum
import dns.name
# Keep pylint quiet on non-windows.
try:
_ = WindowsError # pylint: disable=used-before-assignment
except NameError:
WindowsError = Exception
class ConfigMethod(IntEnum):
Registry = 1
WMI = 2
Win32 = 3
class DnsInfo:
def __init__(self):
self.domain = None
self.nameservers = []
self.search = []
_config_method = ConfigMethod.Registry
if dns._features.have("wmi"):
import threading
import pythoncom # pylint: disable=import-error
import wmi # pylint: disable=import-error
# Prefer WMI by default if wmi is installed.
_config_method = ConfigMethod.WMI
class _WMIGetter(threading.Thread):
# pylint: disable=possibly-used-before-assignment
def __init__(self):
super().__init__()
self.info = DnsInfo()
def run(self):
pythoncom.CoInitialize()
try:
system = wmi.WMI()
for interface in system.Win32_NetworkAdapterConfiguration():
if interface.IPEnabled and interface.DNSServerSearchOrder:
self.info.nameservers = list(interface.DNSServerSearchOrder)
if interface.DNSDomain:
self.info.domain = _config_domain(interface.DNSDomain)
if interface.DNSDomainSuffixSearchOrder:
self.info.search = [
_config_domain(x)
for x in interface.DNSDomainSuffixSearchOrder
]
break
finally:
pythoncom.CoUninitialize()
def get(self):
# We always run in a separate thread to avoid any issues with
# the COM threading model.
self.start()
self.join()
return self.info
else:
class _WMIGetter: # type: ignore
pass
def _config_domain(domain):
# 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("."):
domain = domain[1:]
return dns.name.from_text(domain)
class _RegistryGetter:
def __init__(self):
self.info = DnsInfo()
def _split(self, text):
# The windows registry has used both " " and "," as a delimiter, and while
# it is currently using "," in Windows 10 and later, updates can seemingly
# leave a space in too, e.g. "a, b". So we just convert all commas to
# spaces, and use split() in its default configuration, which splits on
# all whitespace and ignores empty strings.
return text.replace(",", " ").split()
def _config_nameservers(self, nameservers):
for ns in self._split(nameservers):
if ns not in self.info.nameservers:
self.info.nameservers.append(ns)
def _config_search(self, search):
for s in self._split(search):
s = _config_domain(s)
if s not in self.info.search:
self.info.search.append(s)
def _config_fromkey(self, key, always_try_domain):
try:
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")
if dom:
self.info.domain = _config_domain(dom)
except WindowsError:
pass
else:
try:
servers, _ = winreg.QueryValueEx(key, "DhcpNameServer")
except WindowsError:
servers = None
if servers:
self._config_nameservers(servers)
try:
dom, _ = winreg.QueryValueEx(key, "DhcpDomain")
if dom:
self.info.domain = _config_domain(dom)
except WindowsError:
pass
try:
search, _ = winreg.QueryValueEx(key, "SearchList")
except WindowsError:
search = None
if search is None:
try:
search, _ = winreg.QueryValueEx(key, "DhcpSearchList")
except WindowsError:
search = None
if search:
self._config_search(search)
def _is_nic_enabled(self, lm, guid):
# Look in the Windows Registry to determine whether the network
# interface corresponding to the given guid is enabled.
#
# (Code contributed by Paul Marks, thanks!)
#
try:
# This hard-coded location seems to be consistent, at least
# from Windows 2000 through Vista.
connection_key = winreg.OpenKey(
lm,
r"SYSTEM\CurrentControlSet\Control\Network"
r"\{4D36E972-E325-11CE-BFC1-08002BE10318}"
rf"\{guid}\Connection",
)
try:
# The PnpInstanceID points to a key inside Enum
(pnp_id, ttype) = winreg.QueryValueEx(
connection_key, "PnpInstanceID"
)
if ttype != winreg.REG_SZ:
raise ValueError # pragma: no cover
device_key = winreg.OpenKey(
lm, rf"SYSTEM\CurrentControlSet\Enum\{pnp_id}"
)
try:
# Get ConfigFlags for this device
(flags, ttype) = winreg.QueryValueEx(device_key, "ConfigFlags")
if ttype != winreg.REG_DWORD:
raise ValueError # pragma: no cover
# Based on experimentation, bit 0x1 indicates that the
# device is disabled.
#
# XXXRTH I suspect we really want to & with 0x03 so
# that CONFIGFLAGS_REMOVED devices are also ignored,
# but we're shifting to WMI as ConfigFlags is not
# supposed to be used.
return not flags & 0x1
finally:
device_key.Close()
finally:
connection_key.Close()
except Exception: # pragma: no cover
return False
def get(self):
"""Extract resolver configuration from the Windows registry."""
lm = winreg.ConnectRegistry(None, winreg.HKEY_LOCAL_MACHINE)
try:
tcp_params = winreg.OpenKey(
lm, r"SYSTEM\CurrentControlSet\Services\Tcpip\Parameters"
)
try:
self._config_fromkey(tcp_params, True)
finally:
tcp_params.Close()
interfaces = winreg.OpenKey(
lm,
r"SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces",
)
try:
i = 0
while True:
try:
guid = winreg.EnumKey(interfaces, i)
i += 1
key = winreg.OpenKey(interfaces, guid)
try:
if not self._is_nic_enabled(lm, guid):
continue
self._config_fromkey(key, False)
finally:
key.Close()
except OSError:
break
finally:
interfaces.Close()
finally:
lm.Close()
return self.info
class _Win32Getter(_RegistryGetter):
def get(self):
"""Get the attributes using the Windows API."""
# Load the IP Helper library
# # https://learn.microsoft.com/en-us/windows/win32/api/iphlpapi/nf-iphlpapi-getadaptersaddresses
IPHLPAPI = ctypes.WinDLL("Iphlpapi.dll")
# Constants
AF_UNSPEC = 0
ERROR_SUCCESS = 0
GAA_FLAG_INCLUDE_PREFIX = 0x00000010
AF_INET = 2
AF_INET6 = 23
IF_TYPE_SOFTWARE_LOOPBACK = 24
# Define necessary structures
class SOCKADDRV4(ctypes.Structure):
_fields_ = [
("sa_family", wintypes.USHORT),
("sa_data", ctypes.c_ubyte * 14),
]
class SOCKADDRV6(ctypes.Structure):
_fields_ = [
("sa_family", wintypes.USHORT),
("sa_data", ctypes.c_ubyte * 26),
]
class SOCKET_ADDRESS(ctypes.Structure):
_fields_ = [
("lpSockaddr", ctypes.POINTER(SOCKADDRV4)),
("iSockaddrLength", wintypes.INT),
]
class IP_ADAPTER_DNS_SERVER_ADDRESS(ctypes.Structure):
pass # Forward declaration
IP_ADAPTER_DNS_SERVER_ADDRESS._fields_ = [
("Length", wintypes.ULONG),
("Reserved", wintypes.DWORD),
("Next", ctypes.POINTER(IP_ADAPTER_DNS_SERVER_ADDRESS)),
("Address", SOCKET_ADDRESS),
]
class IF_LUID(ctypes.Structure):
_fields_ = [("Value", ctypes.c_ulonglong)]
class NET_IF_NETWORK_GUID(ctypes.Structure):
_fields_ = [("Value", ctypes.c_ubyte * 16)]
class IP_ADAPTER_PREFIX_XP(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_GATEWAY_ADDRESS_LH(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_DNS_SUFFIX(ctypes.Structure):
_fields_ = [
("String", ctypes.c_wchar * 256),
("Next", ctypes.POINTER(ctypes.c_void_p)),
]
class IP_ADAPTER_UNICAST_ADDRESS_LH(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_MULTICAST_ADDRESS_XP(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_ANYCAST_ADDRESS_XP(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_DNS_SERVER_ADDRESS_XP(ctypes.Structure):
pass # Left undefined here for simplicity
class IP_ADAPTER_ADDRESSES(ctypes.Structure):
pass # Forward declaration
IP_ADAPTER_ADDRESSES._fields_ = [
("Length", wintypes.ULONG),
("IfIndex", wintypes.DWORD),
("Next", ctypes.POINTER(IP_ADAPTER_ADDRESSES)),
("AdapterName", ctypes.c_char_p),
("FirstUnicastAddress", ctypes.POINTER(SOCKET_ADDRESS)),
("FirstAnycastAddress", ctypes.POINTER(SOCKET_ADDRESS)),
("FirstMulticastAddress", ctypes.POINTER(SOCKET_ADDRESS)),
(
"FirstDnsServerAddress",
ctypes.POINTER(IP_ADAPTER_DNS_SERVER_ADDRESS),
),
("DnsSuffix", wintypes.LPWSTR),
("Description", wintypes.LPWSTR),
("FriendlyName", wintypes.LPWSTR),
("PhysicalAddress", ctypes.c_ubyte * 8),
("PhysicalAddressLength", wintypes.ULONG),
("Flags", wintypes.ULONG),
("Mtu", wintypes.ULONG),
("IfType", wintypes.ULONG),
("OperStatus", ctypes.c_uint),
# Remaining fields removed for brevity
]
def format_ipv4(sockaddr_in):
return ".".join(map(str, sockaddr_in.sa_data[2:6]))
def format_ipv6(sockaddr_in6):
# The sa_data is:
#
# USHORT sin6_port;
# ULONG sin6_flowinfo;
# IN6_ADDR sin6_addr;
# ULONG sin6_scope_id;
#
# which is 2 + 4 + 16 + 4 = 26 bytes, and we need the plus 6 below
# to be in the sin6_addr range.
parts = [
sockaddr_in6.sa_data[i + 6] << 8 | sockaddr_in6.sa_data[i + 6 + 1]
for i in range(0, 16, 2)
]
return ":".join(f"{part:04x}" for part in parts)
buffer_size = ctypes.c_ulong(15000)
while True:
buffer = ctypes.create_string_buffer(buffer_size.value)
ret_val = IPHLPAPI.GetAdaptersAddresses(
AF_UNSPEC,
GAA_FLAG_INCLUDE_PREFIX,
None,
buffer,
ctypes.byref(buffer_size),
)
if ret_val == ERROR_SUCCESS:
break
elif ret_val != 0x6F: # ERROR_BUFFER_OVERFLOW
print(f"Error retrieving adapter information: {ret_val}")
return
adapter_addresses = ctypes.cast(
buffer, ctypes.POINTER(IP_ADAPTER_ADDRESSES)
)
current_adapter = adapter_addresses
while current_adapter:
# Skip non-operational adapters.
oper_status = current_adapter.contents.OperStatus
if oper_status != 1:
current_adapter = current_adapter.contents.Next
continue
# Exclude loopback adapters.
if current_adapter.contents.IfType == IF_TYPE_SOFTWARE_LOOPBACK:
current_adapter = current_adapter.contents.Next
continue
# Get the domain from the DnsSuffix attribute.
dns_suffix = current_adapter.contents.DnsSuffix
if dns_suffix:
self.info.domain = dns.name.from_text(dns_suffix)
current_dns_server = current_adapter.contents.FirstDnsServerAddress
while current_dns_server:
sockaddr = current_dns_server.contents.Address.lpSockaddr
sockaddr_family = sockaddr.contents.sa_family
ip = None
if sockaddr_family == AF_INET: # IPv4
ip = format_ipv4(sockaddr.contents)
elif sockaddr_family == AF_INET6: # IPv6
sockaddr = ctypes.cast(sockaddr, ctypes.POINTER(SOCKADDRV6))
ip = format_ipv6(sockaddr.contents)
if ip:
if ip not in self.info.nameservers:
self.info.nameservers.append(ip)
current_dns_server = current_dns_server.contents.Next
current_adapter = current_adapter.contents.Next
# Use the registry getter to get the search info, since it is set at the system level.
registry_getter = _RegistryGetter()
info = registry_getter.get()
self.info.search = info.search
return self.info
def set_config_method(method: ConfigMethod) -> None:
global _config_method
_config_method = method
def get_dns_info() -> DnsInfo:
"""Extract resolver configuration."""
if _config_method == ConfigMethod.Win32:
getter = _Win32Getter()
elif _config_method == ConfigMethod.WMI:
getter = _WMIGetter()
else:
getter = _RegistryGetter()
return getter.get()