This commit is contained in:
Mai Gillmann
2019-12-17 14:09:10 +01:00
parent 66e908fc8a
commit 4791d00a43
2122 changed files with 423791 additions and 0 deletions
@@ -0,0 +1,278 @@
# -*- test-case-name: twisted.names.test.test_rfc1982 -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Utilities for handling RFC1982 Serial Number Arithmetic.
@see: U{http://tools.ietf.org/html/rfc1982}
@var RFC4034_TIME_FORMAT: RRSIG Time field presentation format. The Signature
Expiration Time and Inception Time field values MUST be represented either
as an unsigned decimal integer indicating seconds since 1 January 1970
00:00:00 UTC, or in the form YYYYMMDDHHmmSS in UTC. See U{RRSIG Presentation
Format<https://tools.ietf.org/html/rfc4034#section-3.2>}
"""
from __future__ import division, absolute_import
import calendar
from datetime import datetime, timedelta
from twisted.python.compat import nativeString
from twisted.python.util import FancyStrMixin
RFC4034_TIME_FORMAT = '%Y%m%d%H%M%S'
class SerialNumber(FancyStrMixin, object):
"""
An RFC1982 Serial Number.
This class implements RFC1982 DNS Serial Number Arithmetic.
SNA is used in DNS and specifically in DNSSEC as defined in RFC4034 in the
DNSSEC Signature Expiration and Inception Fields.
@see: U{https://tools.ietf.org/html/rfc1982}
@see: U{https://tools.ietf.org/html/rfc4034}
@ivar _serialBits: See C{serialBits} of L{__init__}.
@ivar _number: See C{number} of L{__init__}.
@ivar _modulo: The value at which wrapping will occur.
@ivar _halfRing: Half C{_modulo}. If another L{SerialNumber} value is larger
than this, it would lead to a wrapped value which is larger than the
first and comparisons are therefore ambiguous.
@ivar _maxAdd: Half C{_modulo} plus 1. If another L{SerialNumber} value is
larger than this, it would lead to a wrapped value which is larger than
the first. Comparisons with the original value would therefore be
ambiguous.
"""
showAttributes = (
('_number', 'number', '%d'),
('_serialBits', 'serialBits', '%d'),
)
def __init__(self, number, serialBits=32):
"""
Construct an L{SerialNumber} instance.
@param number: An L{int} which will be stored as the modulo
C{number % 2 ^ serialBits}
@type number: L{int}
@param serialBits: The size of the serial number space. The power of two
which results in one larger than the largest integer corresponding
to a serial number value.
@type serialBits: L{int}
"""
self._serialBits = serialBits
self._modulo = 2 ** serialBits
self._halfRing = 2 ** (serialBits - 1)
self._maxAdd = 2 ** (serialBits - 1) - 1
self._number = int(number) % self._modulo
def _convertOther(self, other):
"""
Check that a foreign object is suitable for use in the comparison or
arithmetic magic methods of this L{SerialNumber} instance. Raise
L{TypeError} if not.
@param other: The foreign L{object} to be checked.
@return: C{other} after compatibility checks and possible coercion.
@raises: L{TypeError} if C{other} is not compatible.
"""
if not isinstance(other, SerialNumber):
raise TypeError(
'cannot compare or combine %r and %r' % (self, other))
if self._serialBits != other._serialBits:
raise TypeError(
'cannot compare or combine SerialNumber instances with '
'different serialBits. %r and %r' % (self, other))
return other
def __str__(self):
"""
Return a string representation of this L{SerialNumber} instance.
@rtype: L{nativeString}
"""
return nativeString('%d' % (self._number,))
def __int__(self):
"""
@return: The integer value of this L{SerialNumber} instance.
@rtype: L{int}
"""
return self._number
def __eq__(self, other):
"""
Allow rich equality comparison with another L{SerialNumber} instance.
@type other: L{SerialNumber}
"""
other = self._convertOther(other)
return other._number == self._number
def __ne__(self, other):
"""
Allow rich equality comparison with another L{SerialNumber} instance.
@type other: L{SerialNumber}
"""
return not self.__eq__(other)
def __lt__(self, other):
"""
Allow I{less than} comparison with another L{SerialNumber} instance.
@type other: L{SerialNumber}
"""
other = self._convertOther(other)
return (
(self._number < other._number
and (other._number - self._number) < self._halfRing)
or
(self._number > other._number
and (self._number - other._number) > self._halfRing)
)
def __gt__(self, other):
"""
Allow I{greater than} comparison with another L{SerialNumber} instance.
@type other: L{SerialNumber}
@rtype: L{bool}
"""
other = self._convertOther(other)
return (
(self._number < other._number
and (other._number - self._number) > self._halfRing)
or
(self._number > other._number
and (self._number - other._number) < self._halfRing)
)
def __le__(self, other):
"""
Allow I{less than or equal} comparison with another L{SerialNumber}
instance.
@type other: L{SerialNumber}
@rtype: L{bool}
"""
other = self._convertOther(other)
return self == other or self < other
def __ge__(self, other):
"""
Allow I{greater than or equal} comparison with another L{SerialNumber}
instance.
@type other: L{SerialNumber}
@rtype: L{bool}
"""
other = self._convertOther(other)
return self == other or self > other
def __add__(self, other):
"""
Allow I{addition} with another L{SerialNumber} instance.
Serial numbers may be incremented by the addition of a positive
integer n, where n is taken from the range of integers
[0 .. (2^(SERIAL_BITS - 1) - 1)]. For a sequence number s, the
result of such an addition, s', is defined as
s' = (s + n) modulo (2 ^ SERIAL_BITS)
where the addition and modulus operations here act upon values that are
non-negative values of unbounded size in the usual ways of integer
arithmetic.
Addition of a value outside the range
[0 .. (2^(SERIAL_BITS - 1) - 1)] is undefined.
@see: U{http://tools.ietf.org/html/rfc1982#section-3.1}
@type other: L{SerialNumber}
@rtype: L{SerialNumber}
@raises: L{ArithmeticError} if C{other} is more than C{_maxAdd}
ie more than half the maximum value of this serial number.
"""
other = self._convertOther(other)
if other._number <= self._maxAdd:
return SerialNumber(
(self._number + other._number) % self._modulo,
serialBits=self._serialBits)
else:
raise ArithmeticError(
'value %r outside the range 0 .. %r' % (
other._number, self._maxAdd,))
def __hash__(self):
"""
Allow L{SerialNumber} instances to be hashed for use as L{dict} keys.
@rtype: L{int}
"""
return hash(self._number)
@classmethod
def fromRFC4034DateString(cls, utcDateString):
"""
Create an L{SerialNumber} instance from a date string in format
'YYYYMMDDHHMMSS' described in U{RFC4034
3.2<https://tools.ietf.org/html/rfc4034#section-3.2>}.
The L{SerialNumber} instance stores the date as a 32bit UNIX timestamp.
@see: U{https://tools.ietf.org/html/rfc4034#section-3.1.5}
@param utcDateString: A UTC date/time string of format I{YYMMDDhhmmss}
which will be converted to seconds since the UNIX epoch.
@type utcDateString: L{unicode}
@return: An L{SerialNumber} instance containing the supplied date as a
32bit UNIX timestamp.
"""
parsedDate = datetime.strptime(utcDateString, RFC4034_TIME_FORMAT)
secondsSinceEpoch = calendar.timegm(parsedDate.utctimetuple())
return cls(secondsSinceEpoch, serialBits=32)
def toRFC4034DateString(self):
"""
Calculate a date by treating the current L{SerialNumber} value as a UNIX
timestamp and return a date string in the format described in
U{RFC4034 3.2<https://tools.ietf.org/html/rfc4034#section-3.2>}.
@return: The date string.
"""
# Can't use datetime.utcfromtimestamp, because it seems to overflow the
# signed 32bit int used in the underlying C library. SNA is unsigned
# and capable of handling all timestamps up to 2**32.
d = datetime(1970, 1, 1) + timedelta(seconds=self._number)
return nativeString(d.strftime(RFC4034_TIME_FORMAT))
__all__ = ['SerialNumber']
@@ -0,0 +1,99 @@
# -*- test-case-name: twisted.names.test.test_resolve -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Lookup a name using multiple resolvers.
Future Plans: This needs someway to specify which resolver answered
the query, or someway to specify (authority|ttl|cache behavior|more?)
"""
from __future__ import division, absolute_import
from zope.interface import implementer
from twisted.internet import defer, interfaces
from twisted.names import dns, common, error
class FailureHandler:
def __init__(self, resolver, query, timeout):
self.resolver = resolver
self.query = query
self.timeout = timeout
def __call__(self, failure):
# AuthoritativeDomainErrors should halt resolution attempts
failure.trap(dns.DomainError, defer.TimeoutError, NotImplementedError)
return self.resolver(self.query, self.timeout)
@implementer(interfaces.IResolver)
class ResolverChain(common.ResolverBase):
"""
Lookup an address using multiple L{IResolver}s
"""
def __init__(self, resolvers):
"""
@type resolvers: L{list}
@param resolvers: A L{list} of L{IResolver} providers.
"""
common.ResolverBase.__init__(self)
self.resolvers = resolvers
def _lookup(self, name, cls, type, timeout):
"""
Build a L{dns.Query} for the given parameters and dispatch it
to each L{IResolver} in C{self.resolvers} until an answer or
L{error.AuthoritativeDomainError} is returned.
@type name: C{str}
@param name: DNS name to resolve.
@type type: C{int}
@param type: DNS record type.
@type cls: C{int}
@param cls: DNS record class.
@type timeout: Sequence of C{int}
@param timeout: Number of seconds after which to reissue the query.
When the last timeout expires, the query is considered failed.
@rtype: L{Deferred}
@return: A L{Deferred} which fires with a three-tuple of lists of
L{twisted.names.dns.RRHeader} instances. The first element of the
tuple gives answers. The second element of the tuple gives
authorities. The third element of the tuple gives additional
information. The L{Deferred} may instead fail with one of the
exceptions defined in L{twisted.names.error} or with
C{NotImplementedError}.
"""
if not self.resolvers:
return defer.fail(error.DomainError())
q = dns.Query(name, type, cls)
d = self.resolvers[0].query(q, timeout)
for r in self.resolvers[1:]:
d = d.addErrback(
FailureHandler(r.query, q, timeout)
)
return d
def lookupAllRecords(self, name, timeout=None):
# XXX: Why is this necessary? dns.ALL_RECORDS queries should
# be handled just the same as any other type by _lookup
# above. If I remove this method all names tests still
# pass. See #6604 -rwall
if not self.resolvers:
return defer.fail(error.DomainError())
d = self.resolvers[0].lookupAllRecords(name, timeout)
for r in self.resolvers[1:]:
d = d.addErrback(
FailureHandler(r.lookupAllRecords, name, timeout)
)
return d
@@ -0,0 +1,590 @@
# -*- test-case-name: twisted.names.test.test_names,twisted.names.test.test_server -*-
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Async DNS server
Future plans:
- Better config file format maybe
- Make sure to differentiate between different classes
- notice truncation bit
Important: No additional processing is done on some of the record types.
This violates the most basic RFC and is just plain annoying
for resolvers to deal with. Fix it.
@author: Jp Calderone
"""
from __future__ import division, absolute_import
import time
from twisted.internet import protocol
from twisted.names import dns, resolve
from twisted.python import log
class DNSServerFactory(protocol.ServerFactory):
"""
Server factory and tracker for L{DNSProtocol} connections. This class also
provides records for responses to DNS queries.
@ivar cache: A L{Cache<twisted.names.cache.CacheResolver>} instance whose
C{cacheResult} method is called when a response is received from one of
C{clients}. Defaults to L{None} if no caches are specified. See
C{caches} of L{__init__} for more details.
@type cache: L{Cache<twisted.names.cache.CacheResolver>} or L{None}
@ivar canRecurse: A flag indicating whether this server is capable of
performing recursive DNS resolution.
@type canRecurse: L{bool}
@ivar resolver: A L{resolve.ResolverChain} containing an ordered list of
C{authorities}, C{caches} and C{clients} to which queries will be
dispatched.
@type resolver: L{resolve.ResolverChain}
@ivar verbose: See L{__init__}
@ivar connections: A list of all the connected L{DNSProtocol} instances
using this object as their controller.
@type connections: C{list} of L{DNSProtocol} instances
@ivar protocol: A callable used for building a DNS stream protocol. Called
by L{DNSServerFactory.buildProtocol} and passed the L{DNSServerFactory}
instance as the one and only positional argument. Defaults to
L{dns.DNSProtocol}.
@type protocol: L{IProtocolFactory} constructor
@ivar _messageFactory: A response message constructor with an initializer
signature matching L{dns.Message.__init__}.
@type _messageFactory: C{callable}
"""
protocol = dns.DNSProtocol
cache = None
_messageFactory = dns.Message
def __init__(self, authorities=None, caches=None, clients=None, verbose=0):
"""
@param authorities: Resolvers which provide authoritative answers.
@type authorities: L{list} of L{IResolver} providers
@param caches: Resolvers which provide cached non-authoritative
answers. The first cache instance is assigned to
C{DNSServerFactory.cache} and its C{cacheResult} method will be
called when a response is received from one of C{clients}.
@type caches: L{list} of L{Cache<twisted.names.cache.CacheResolver>} instances
@param clients: Resolvers which are capable of performing recursive DNS
lookups.
@type clients: L{list} of L{IResolver} providers
@param verbose: An integer controlling the verbosity of logging of
queries and responses. Default is C{0} which means no logging. Set
to C{2} to enable logging of full query and response messages.
@type verbose: L{int}
"""
resolvers = []
if authorities is not None:
resolvers.extend(authorities)
if caches is not None:
resolvers.extend(caches)
if clients is not None:
resolvers.extend(clients)
self.canRecurse = not not clients
self.resolver = resolve.ResolverChain(resolvers)
self.verbose = verbose
if caches:
self.cache = caches[-1]
self.connections = []
def _verboseLog(self, *args, **kwargs):
"""
Log a message only if verbose logging is enabled.
@param args: Positional arguments which will be passed to C{log.msg}
@param kwargs: Keyword arguments which will be passed to C{log.msg}
"""
if self.verbose > 0:
log.msg(*args, **kwargs)
def buildProtocol(self, addr):
p = self.protocol(self)
p.factory = self
return p
def connectionMade(self, protocol):
"""
Track a newly connected L{DNSProtocol}.
@param protocol: The protocol instance to be tracked.
@type protocol: L{dns.DNSProtocol}
"""
self.connections.append(protocol)
def connectionLost(self, protocol):
"""
Stop tracking a no-longer connected L{DNSProtocol}.
@param protocol: The tracked protocol instance to be which has been
lost.
@type protocol: L{dns.DNSProtocol}
"""
self.connections.remove(protocol)
def sendReply(self, protocol, message, address):
"""
Send a response C{message} to a given C{address} via the supplied
C{protocol}.
Message payload will be logged if C{DNSServerFactory.verbose} is C{>1}.
@param protocol: The DNS protocol instance to which to send the message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The DNS message to be sent.
@type message: L{dns.Message}
@param address: The address to which the message will be sent or L{None}
if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
if self.verbose > 1:
s = ' '.join([str(a.payload) for a in message.answers])
auth = ' '.join([str(a.payload) for a in message.authority])
add = ' '.join([str(a.payload) for a in message.additional])
if not s:
log.msg("Replying with no answers")
else:
log.msg("Answers are " + s)
log.msg("Authority is " + auth)
log.msg("Additional is " + add)
if address is None:
protocol.writeMessage(message)
else:
protocol.writeMessage(message, address)
self._verboseLog(
"Processed query in %0.3f seconds" % (
time.time() - message.timeReceived))
def _responseFromMessage(self, message, rCode=dns.OK,
answers=None, authority=None, additional=None):
"""
Generate a L{Message} instance suitable for use as the response to
C{message}.
C{queries} will be copied from the request to the response.
C{rCode}, C{answers}, C{authority} and C{additional} will be assigned to
the response, if supplied.
The C{recAv} flag will be set on the response if the C{canRecurse} flag
on this L{DNSServerFactory} is set to L{True}.
The C{auth} flag will be set on the response if *any* of the supplied
C{answers} have their C{auth} flag set to L{True}.
The response will have the same C{maxSize} as the request.
Additionally, the response will have a C{timeReceived} attribute whose
value is that of the original request and the
@see: L{dns._responseFromMessage}
@param message: The request message
@type message: L{Message}
@param rCode: The response code which will be assigned to the response.
@type message: L{int}
@param answers: An optional list of answer records which will be
assigned to the response.
@type answers: L{list} of L{dns.RRHeader}
@param authority: An optional list of authority records which will be
assigned to the response.
@type authority: L{list} of L{dns.RRHeader}
@param additional: An optional list of additional records which will be
assigned to the response.
@type additional: L{list} of L{dns.RRHeader}
@return: A response L{Message} instance.
@rtype: L{Message}
"""
if answers is None:
answers = []
if authority is None:
authority = []
if additional is None:
additional = []
authoritativeAnswer = False
for x in answers:
if x.isAuthoritative():
authoritativeAnswer = True
break
response = dns._responseFromMessage(
responseConstructor=self._messageFactory,
message=message,
recAv=self.canRecurse,
rCode=rCode,
auth=authoritativeAnswer
)
# XXX: Timereceived is a hack which probably shouldn't be tacked onto
# the message. Use getattr here so that we don't have to set the
# timereceived on every message in the tests. See #6957.
response.timeReceived = getattr(message, 'timeReceived', None)
# XXX: This is another hack. dns.Message.decode sets maxSize=0 which
# means that responses are never truncated. I'll maintain that behaviour
# here until #6949 is resolved.
response.maxSize = message.maxSize
response.answers = answers
response.authority = authority
response.additional = additional
return response
def gotResolverResponse(self, response, protocol, message, address):
"""
A callback used by L{DNSServerFactory.handleQuery} for handling the
deferred response from C{self.resolver.query}.
Constructs a response message by combining the original query message
with the resolved answer, authority and additional records.
Marks the response message as authoritative if any of the resolved
answers are found to be authoritative.
The resolved answers count will be logged if C{DNSServerFactory.verbose}
is C{>1}.
@param response: Answer records, authority records and additional records
@type response: L{tuple} of L{list} of L{dns.RRHeader} instances
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
ans, auth, add = response
response = self._responseFromMessage(
message=message, rCode=dns.OK,
answers=ans, authority=auth, additional=add)
self.sendReply(protocol, response, address)
l = len(ans) + len(auth) + len(add)
self._verboseLog("Lookup found %d record%s" % (l, l != 1 and "s" or ""))
if self.cache and l:
self.cache.cacheResult(
message.queries[0], (ans, auth, add)
)
def gotResolverError(self, failure, protocol, message, address):
"""
A callback used by L{DNSServerFactory.handleQuery} for handling deferred
errors from C{self.resolver.query}.
Constructs a response message from the original query message by
assigning a suitable error code to C{rCode}.
An error message will be logged if C{DNSServerFactory.verbose} is C{>1}.
@param failure: The reason for the failed resolution (as reported by
C{self.resolver.query}).
@type failure: L{Failure<twisted.python.failure.Failure>}
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
if failure.check(dns.DomainError, dns.AuthoritativeDomainError):
rCode = dns.ENAME
else:
rCode = dns.ESERVER
log.err(failure)
response = self._responseFromMessage(message=message, rCode=rCode)
self.sendReply(protocol, response, address)
self._verboseLog("Lookup failed")
def handleQuery(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} when a query message is
received.
Takes the first query from the received message and dispatches it to
C{self.resolver.query}.
Adds callbacks L{DNSServerFactory.gotResolverResponse} and
L{DNSServerFactory.gotResolverError} to the resulting deferred.
Note: Multiple queries in a single message are not supported because
there is no standard way to respond with multiple rCodes, auth,
etc. This is consistent with other DNS server implementations. See
U{http://tools.ietf.org/html/draft-ietf-dnsext-edns1-03} for a proposed
solution.
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
@return: A C{deferred} which fires with the resolved result or error of
the first query in C{message}.
@rtype: L{Deferred<twisted.internet.defer.Deferred>}
"""
query = message.queries[0]
return self.resolver.query(query).addCallback(
self.gotResolverResponse, protocol, message, address
).addErrback(
self.gotResolverError, protocol, message, address
)
def handleInverseQuery(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} when an inverse query
message is received.
Replies with a I{Not Implemented} error by default.
An error message will be logged if C{DNSServerFactory.verbose} is C{>1}.
Override in a subclass.
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
message.rCode = dns.ENOTIMP
self.sendReply(protocol, message, address)
self._verboseLog("Inverse query from %r" % (address,))
def handleStatus(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} when a status message is
received.
Replies with a I{Not Implemented} error by default.
An error message will be logged if C{DNSServerFactory.verbose} is C{>1}.
Override in a subclass.
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
message.rCode = dns.ENOTIMP
self.sendReply(protocol, message, address)
self._verboseLog("Status request from %r" % (address,))
def handleNotify(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} when a notify message is
received.
Replies with a I{Not Implemented} error by default.
An error message will be logged if C{DNSServerFactory.verbose} is C{>1}.
Override in a subclass.
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
message.rCode = dns.ENOTIMP
self.sendReply(protocol, message, address)
self._verboseLog("Notify message from %r" % (address,))
def handleOther(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} when a message with
unrecognised I{OPCODE} is received.
Replies with a I{Not Implemented} error by default.
An error message will be logged if C{DNSServerFactory.verbose} is C{>1}.
Override in a subclass.
@param protocol: The DNS protocol instance to which to send a response
message.
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param message: The original DNS query message for which a response
message will be constructed.
@type message: L{dns.Message}
@param address: The address to which the response message will be sent
or L{None} if C{protocol} is a stream protocol.
@type address: L{tuple} or L{None}
"""
message.rCode = dns.ENOTIMP
self.sendReply(protocol, message, address)
self._verboseLog(
"Unknown op code (%d) from %r" % (message.opCode, address))
def messageReceived(self, message, proto, address=None):
"""
L{DNSServerFactory.messageReceived} is called by protocols which are
under the control of this L{DNSServerFactory} whenever they receive a
DNS query message or an unexpected / duplicate / late DNS response
message.
L{DNSServerFactory.allowQuery} is called with the received message,
protocol and origin address. If it returns L{False}, a C{dns.EREFUSED}
response is sent back to the client.
Otherwise the received message is dispatched to one of
L{DNSServerFactory.handleQuery}, L{DNSServerFactory.handleInverseQuery},
L{DNSServerFactory.handleStatus}, L{DNSServerFactory.handleNotify}, or
L{DNSServerFactory.handleOther} depending on the I{OPCODE} of the
received message.
If C{DNSServerFactory.verbose} is C{>0} all received messages will be
logged in more or less detail depending on the value of C{verbose}.
@param message: The DNS message that was received.
@type message: L{dns.Message}
@param proto: The DNS protocol instance which received the message
@type proto: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param address: The address from which the message was received. Only
provided for messages received by datagram protocols. The origin of
Messages received from stream protocols can be gleaned from the
protocol C{transport} attribute.
@type address: L{tuple} or L{None}
"""
message.timeReceived = time.time()
if self.verbose:
if self.verbose > 1:
s = ' '.join([str(q) for q in message.queries])
else:
s = ' '.join([dns.QUERY_TYPES.get(q.type, 'UNKNOWN')
for q in message.queries])
if not len(s):
log.msg(
"Empty query from %r" % (
(address or proto.transport.getPeer()),))
else:
log.msg(
"%s query from %r" % (
s, address or proto.transport.getPeer()))
if not self.allowQuery(message, proto, address):
message.rCode = dns.EREFUSED
self.sendReply(proto, message, address)
elif message.opCode == dns.OP_QUERY:
self.handleQuery(message, proto, address)
elif message.opCode == dns.OP_INVERSE:
self.handleInverseQuery(message, proto, address)
elif message.opCode == dns.OP_STATUS:
self.handleStatus(message, proto, address)
elif message.opCode == dns.OP_NOTIFY:
self.handleNotify(message, proto, address)
else:
self.handleOther(message, proto, address)
def allowQuery(self, message, protocol, address):
"""
Called by L{DNSServerFactory.messageReceived} to decide whether to
process a received message or to reply with C{dns.EREFUSED}.
This default implementation permits anything but empty queries.
Override in a subclass to implement alternative policies.
@param message: The DNS message that was received.
@type message: L{dns.Message}
@param protocol: The DNS protocol instance which received the message
@type protocol: L{dns.DNSDatagramProtocol} or L{dns.DNSProtocol}
@param address: The address from which the message was received. Only
provided for messages received by datagram protocols. The origin of
Messages received from stream protocols can be gleaned from the
protocol C{transport} attribute.
@type address: L{tuple} or L{None}
@return: L{True} if the received message contained one or more queries,
else L{False}.
@rtype: L{bool}
"""
return len(message.queries)
@@ -0,0 +1,133 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Tests for L{twisted.names.common}.
"""
from __future__ import division, absolute_import
from zope.interface.verify import verifyClass
from twisted.internet.interfaces import IResolver
from twisted.trial.unittest import SynchronousTestCase
from twisted.python.failure import Failure
from twisted.names.common import ResolverBase
from twisted.names.dns import EFORMAT, ESERVER, ENAME, ENOTIMP, EREFUSED, Query
from twisted.names.error import DNSFormatError, DNSServerError, DNSNameError
from twisted.names.error import DNSNotImplementedError, DNSQueryRefusedError
from twisted.names.error import DNSUnknownError
class ExceptionForCodeTests(SynchronousTestCase):
"""
Tests for L{ResolverBase.exceptionForCode}.
"""
def setUp(self):
self.exceptionForCode = ResolverBase().exceptionForCode
def test_eformat(self):
"""
L{ResolverBase.exceptionForCode} converts L{EFORMAT} to
L{DNSFormatError}.
"""
self.assertIs(self.exceptionForCode(EFORMAT), DNSFormatError)
def test_eserver(self):
"""
L{ResolverBase.exceptionForCode} converts L{ESERVER} to
L{DNSServerError}.
"""
self.assertIs(self.exceptionForCode(ESERVER), DNSServerError)
def test_ename(self):
"""
L{ResolverBase.exceptionForCode} converts L{ENAME} to L{DNSNameError}.
"""
self.assertIs(self.exceptionForCode(ENAME), DNSNameError)
def test_enotimp(self):
"""
L{ResolverBase.exceptionForCode} converts L{ENOTIMP} to
L{DNSNotImplementedError}.
"""
self.assertIs(self.exceptionForCode(ENOTIMP), DNSNotImplementedError)
def test_erefused(self):
"""
L{ResolverBase.exceptionForCode} converts L{EREFUSED} to
L{DNSQueryRefusedError}.
"""
self.assertIs(self.exceptionForCode(EREFUSED), DNSQueryRefusedError)
def test_other(self):
"""
L{ResolverBase.exceptionForCode} converts any other response code to
L{DNSUnknownError}.
"""
self.assertIs(self.exceptionForCode(object()), DNSUnknownError)
class QueryTests(SynchronousTestCase):
"""
Tests for L{ResolverBase.query}.
"""
def test_resolverBaseProvidesIResolver(self):
"""
L{ResolverBase} provides the L{IResolver} interface.
"""
verifyClass(IResolver, ResolverBase)
def test_typeToMethodDispatch(self):
"""
L{ResolverBase.query} looks up a method to invoke using the type of the
query passed to it and the C{typeToMethod} mapping on itself.
"""
results = []
resolver = ResolverBase()
resolver.typeToMethod = {
12345: lambda query, timeout: results.append((query, timeout))}
query = Query(name=b"example.com", type=12345)
resolver.query(query, 123)
self.assertEqual([(b"example.com", 123)], results)
def test_typeToMethodResult(self):
"""
L{ResolverBase.query} returns a L{Deferred} which fires with the result
of the method found in the C{typeToMethod} mapping for the type of the
query passed to it.
"""
expected = object()
resolver = ResolverBase()
resolver.typeToMethod = {54321: lambda query, timeout: expected}
query = Query(name=b"example.com", type=54321)
queryDeferred = resolver.query(query, 123)
result = []
queryDeferred.addBoth(result.append)
self.assertEqual(expected, result[0])
def test_unknownQueryType(self):
"""
L{ResolverBase.query} returns a L{Deferred} which fails with
L{NotImplementedError} when called with a query of a type not present in
its C{typeToMethod} dictionary.
"""
resolver = ResolverBase()
resolver.typeToMethod = {}
query = Query(name=b"example.com", type=12345)
queryDeferred = resolver.query(query, 123)
result = []
queryDeferred.addBoth(result.append)
self.assertIsInstance(result[0], Failure)
result[0].trap(NotImplementedError)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,444 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test cases for L{twisted.names.rfc1982}.
"""
from __future__ import division, absolute_import
import calendar
from datetime import datetime
from functools import partial
from twisted.names._rfc1982 import SerialNumber
from twisted.trial import unittest
class SerialNumberTests(unittest.TestCase):
"""
Tests for L{SerialNumber}.
"""
def test_serialBitsDefault(self):
"""
L{SerialNumber.serialBits} has default value 32.
"""
self.assertEqual(SerialNumber(1)._serialBits, 32)
def test_serialBitsOverride(self):
"""
L{SerialNumber.__init__} accepts a C{serialBits} argument whose value is
assigned to L{SerialNumber.serialBits}.
"""
self.assertEqual(SerialNumber(1, serialBits=8)._serialBits, 8)
def test_repr(self):
"""
L{SerialNumber.__repr__} returns a string containing number and
serialBits.
"""
self.assertEqual(
'<SerialNumber number=123 serialBits=32>',
repr(SerialNumber(123, serialBits=32))
)
def test_str(self):
"""
L{SerialNumber.__str__} returns a string representation of the current
value.
"""
self.assertEqual(str(SerialNumber(123)), '123')
def test_int(self):
"""
L{SerialNumber.__int__} returns an integer representation of the current
value.
"""
self.assertEqual(int(SerialNumber(123)), 123)
def test_hash(self):
"""
L{SerialNumber.__hash__} allows L{SerialNumber} instances to be hashed
for use as dictionary keys.
"""
self.assertEqual(hash(SerialNumber(1)), hash(SerialNumber(1)))
self.assertNotEqual(hash(SerialNumber(1)), hash(SerialNumber(2)))
def test_convertOtherSerialBitsMismatch(self):
"""
L{SerialNumber._convertOther} raises L{TypeError} if the other
SerialNumber instance has a different C{serialBits} value.
"""
s1 = SerialNumber(0, serialBits=8)
s2 = SerialNumber(0, serialBits=16)
self.assertRaises(
TypeError,
s1._convertOther,
s2
)
def test_eq(self):
"""
L{SerialNumber.__eq__} provides rich equality comparison.
"""
self.assertEqual(SerialNumber(1), SerialNumber(1))
def test_eqForeignType(self):
"""
== comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) == object())
def test_ne(self):
"""
L{SerialNumber.__ne__} provides rich equality comparison.
"""
self.assertFalse(SerialNumber(1) != SerialNumber(1))
self.assertNotEqual(SerialNumber(1), SerialNumber(2))
def test_neForeignType(self):
"""
!= comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) != object())
def test_le(self):
"""
L{SerialNumber.__le__} provides rich <= comparison.
"""
self.assertTrue(SerialNumber(1) <= SerialNumber(1))
self.assertTrue(SerialNumber(1) <= SerialNumber(2))
def test_leForeignType(self):
"""
<= comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) <= object())
def test_ge(self):
"""
L{SerialNumber.__ge__} provides rich >= comparison.
"""
self.assertTrue(SerialNumber(1) >= SerialNumber(1))
self.assertTrue(SerialNumber(2) >= SerialNumber(1))
def test_geForeignType(self):
"""
>= comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) >= object())
def test_lt(self):
"""
L{SerialNumber.__lt__} provides rich < comparison.
"""
self.assertTrue(SerialNumber(1) < SerialNumber(2))
def test_ltForeignType(self):
"""
< comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) < object())
def test_gt(self):
"""
L{SerialNumber.__gt__} provides rich > comparison.
"""
self.assertTrue(SerialNumber(2) > SerialNumber(1))
def test_gtForeignType(self):
"""
> comparison of L{SerialNumber} with a non-L{SerialNumber} instance
raises L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(2) > object())
def test_add(self):
"""
L{SerialNumber.__add__} allows L{SerialNumber} instances to be summed.
"""
self.assertEqual(SerialNumber(1) + SerialNumber(1), SerialNumber(2))
def test_addForeignType(self):
"""
Addition of L{SerialNumber} with a non-L{SerialNumber} instance raises
L{TypeError}.
"""
self.assertRaises(TypeError, lambda: SerialNumber(1) + object())
def test_addOutOfRangeHigh(self):
"""
L{SerialNumber} cannot be added with other SerialNumber values larger
than C{_maxAdd}.
"""
maxAdd = SerialNumber(1)._maxAdd
self.assertRaises(
ArithmeticError,
lambda: SerialNumber(1) + SerialNumber(maxAdd + 1))
def test_maxVal(self):
"""
L{SerialNumber.__add__} returns a wrapped value when s1 plus the s2
would result in a value greater than the C{maxVal}.
"""
s = SerialNumber(1)
maxVal = s._halfRing + s._halfRing - 1
maxValPlus1 = maxVal + 1
self.assertTrue(SerialNumber(maxValPlus1) > SerialNumber(maxVal))
self.assertEqual(SerialNumber(maxValPlus1), SerialNumber(0))
def test_fromRFC4034DateString(self):
"""
L{SerialNumber.fromRFC4034DateString} accepts a datetime string argument
of the form 'YYYYMMDDhhmmss' and returns an L{SerialNumber} instance
whose value is the unix timestamp corresponding to that UTC date.
"""
self.assertEqual(
SerialNumber(1325376000),
SerialNumber.fromRFC4034DateString('20120101000000')
)
def test_toRFC4034DateString(self):
"""
L{DateSerialNumber.toRFC4034DateString} interprets the current value as
a unix timestamp and returns a date string representation of that date.
"""
self.assertEqual(
'20120101000000',
SerialNumber(1325376000).toRFC4034DateString()
)
def test_unixEpoch(self):
"""
L{SerialNumber.toRFC4034DateString} stores 32bit timestamps relative to
the UNIX epoch.
"""
self.assertEqual(
SerialNumber(0).toRFC4034DateString(),
'19700101000000'
)
def test_Y2106Problem(self):
"""
L{SerialNumber} wraps unix timestamps in the year 2106.
"""
self.assertEqual(
SerialNumber(-1).toRFC4034DateString(),
'21060207062815'
)
def test_Y2038Problem(self):
"""
L{SerialNumber} raises ArithmeticError when used to add dates more than
68 years in the future.
"""
maxAddTime = calendar.timegm(
datetime(2038, 1, 19, 3, 14, 7).utctimetuple())
self.assertEqual(
maxAddTime,
SerialNumber(0)._maxAdd,
)
self.assertRaises(
ArithmeticError,
lambda: SerialNumber(0) + SerialNumber(maxAddTime + 1))
def assertUndefinedComparison(testCase, s1, s2):
"""
A custom assertion for L{SerialNumber} values that cannot be meaningfully
compared.
"Note that there are some pairs of values s1 and s2 for which s1 is not
equal to s2, but for which s1 is neither greater than, nor less than, s2.
An attempt to use these ordering operators on such pairs of values produces
an undefined result."
@see: U{https://tools.ietf.org/html/rfc1982#section-3.2}
@param testCase: The L{unittest.TestCase} on which to call assertion
methods.
@type testCase: L{unittest.TestCase}
@param s1: The first value to compare.
@type s1: L{SerialNumber}
@param s2: The second value to compare.
@type s2: L{SerialNumber}
"""
testCase.assertFalse(s1 == s2)
testCase.assertFalse(s1 <= s2)
testCase.assertFalse(s1 < s2)
testCase.assertFalse(s1 > s2)
testCase.assertFalse(s1 >= s2)
serialNumber2 = partial(SerialNumber, serialBits=2)
class SerialNumber2BitTests(unittest.TestCase):
"""
Tests for correct answers to example calculations in RFC1982 5.1.
The simplest meaningful serial number space has SERIAL_BITS == 2. In this
space, the integers that make up the serial number space are 0, 1, 2, and 3.
That is, 3 == 2^SERIAL_BITS - 1.
https://tools.ietf.org/html/rfc1982#section-5.1
"""
def test_maxadd(self):
"""
In this space, the largest integer that it is meaningful to add to a
sequence number is 2^(SERIAL_BITS - 1) - 1, or 1.
"""
self.assertEqual(SerialNumber(0, serialBits=2)._maxAdd, 1)
def test_add(self):
"""
Then, as defined 0+1 == 1, 1+1 == 2, 2+1 == 3, and 3+1 == 0.
"""
self.assertEqual(serialNumber2(0) + serialNumber2(1), serialNumber2(1))
self.assertEqual(serialNumber2(1) + serialNumber2(1), serialNumber2(2))
self.assertEqual(serialNumber2(2) + serialNumber2(1), serialNumber2(3))
self.assertEqual(serialNumber2(3) + serialNumber2(1), serialNumber2(0))
def test_gt(self):
"""
Further, 1 > 0, 2 > 1, 3 > 2, and 0 > 3.
"""
self.assertTrue(serialNumber2(1) > serialNumber2(0))
self.assertTrue(serialNumber2(2) > serialNumber2(1))
self.assertTrue(serialNumber2(3) > serialNumber2(2))
self.assertTrue(serialNumber2(0) > serialNumber2(3))
def test_undefined(self):
"""
It is undefined whether 2 > 0 or 0 > 2, and whether 1 > 3 or 3 > 1.
"""
assertUndefinedComparison(self, serialNumber2(2), serialNumber2(0))
assertUndefinedComparison(self, serialNumber2(0), serialNumber2(2))
assertUndefinedComparison(self, serialNumber2(1), serialNumber2(3))
assertUndefinedComparison(self, serialNumber2(3), serialNumber2(1))
serialNumber8 = partial(SerialNumber, serialBits=8)
class SerialNumber8BitTests(unittest.TestCase):
"""
Tests for correct answers to example calculations in RFC1982 5.2.
Consider the case where SERIAL_BITS == 8. In this space the integers that
make up the serial number space are 0, 1, 2, ... 254, 255. 255 ==
2^SERIAL_BITS - 1.
https://tools.ietf.org/html/rfc1982#section-5.2
"""
def test_maxadd(self):
"""
In this space, the largest integer that it is meaningful to add to a
sequence number is 2^(SERIAL_BITS - 1) - 1, or 127.
"""
self.assertEqual(SerialNumber(0, serialBits=8)._maxAdd, 127)
def test_add(self):
"""
Addition is as expected in this space, for example: 255+1 == 0,
100+100 == 200, and 200+100 == 44.
"""
self.assertEqual(
serialNumber8(255) + serialNumber8(1), serialNumber8(0))
self.assertEqual(
serialNumber8(100) + serialNumber8(100), serialNumber8(200))
self.assertEqual(
serialNumber8(200) + serialNumber8(100), serialNumber8(44))
def test_gt(self):
"""
Comparison is more interesting, 1 > 0, 44 > 0, 100 > 0, 100 > 44,
200 > 100, 255 > 200, 0 > 255, 100 > 255, 0 > 200, and 44 > 200.
"""
self.assertTrue(serialNumber8(1) > serialNumber8(0))
self.assertTrue(serialNumber8(44) > serialNumber8(0))
self.assertTrue(serialNumber8(100) > serialNumber8(0))
self.assertTrue(serialNumber8(100) > serialNumber8(44))
self.assertTrue(serialNumber8(200) > serialNumber8(100))
self.assertTrue(serialNumber8(255) > serialNumber8(200))
self.assertTrue(serialNumber8(100) > serialNumber8(255))
self.assertTrue(serialNumber8(0) > serialNumber8(200))
self.assertTrue(serialNumber8(44) > serialNumber8(200))
def test_surprisingAddition(self):
"""
Note that 100+100 > 100, but that (100+100)+100 < 100. Incrementing a
serial number can cause it to become "smaller". Of course, incrementing
by a smaller number will allow many more increments to be made before
this occurs. However this is always something to be aware of, it can
cause surprising errors, or be useful as it is the only defined way to
actually cause a serial number to decrease.
"""
self.assertTrue(
serialNumber8(100) + serialNumber8(100) > serialNumber8(100))
self.assertTrue(
serialNumber8(100) + serialNumber8(100) + serialNumber8(100)
< serialNumber8(100))
def test_undefined(self):
"""
The pairs of values 0 and 128, 1 and 129, 2 and 130, etc, to 127 and 255
are not equal, but in each pair, neither number is defined as being
greater than, or less than, the other.
"""
assertUndefinedComparison(self, serialNumber8(0), serialNumber8(128))
assertUndefinedComparison(self, serialNumber8(1), serialNumber8(129))
assertUndefinedComparison(self, serialNumber8(2), serialNumber8(130))
assertUndefinedComparison(self, serialNumber8(127), serialNumber8(255))
@@ -0,0 +1,734 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Test cases for Twisted.names' root resolver.
"""
from zope.interface import implementer
from zope.interface.verify import verifyClass
from twisted.python.log import msg
from twisted.trial import util
from twisted.trial.unittest import SynchronousTestCase, TestCase
from twisted.internet.defer import Deferred, succeed, gatherResults, TimeoutError
from twisted.internet.interfaces import IResolverSimple
from twisted.names import client, root
from twisted.names.root import Resolver
from twisted.names.dns import (
IN, HS, A, NS, CNAME, OK, ENAME, Record_CNAME,
Name, Query, Message, RRHeader, Record_A, Record_NS)
from twisted.names.error import DNSNameError, ResolverError
from twisted.names.test.test_util import MemoryReactor
def getOnePayload(results):
"""
From the result of a L{Deferred} returned by L{IResolver.lookupAddress},
return the payload of the first record in the answer section.
"""
ans, auth, add = results
return ans[0].payload
def getOneAddress(results):
"""
From the result of a L{Deferred} returned by L{IResolver.lookupAddress},
return the first IPv4 address from the answer section.
"""
return getOnePayload(results).dottedQuad()
class RootResolverTests(TestCase):
"""
Tests for L{twisted.names.root.Resolver}.
"""
def _queryTest(self, filter):
"""
Invoke L{Resolver._query} and verify that it sends the correct DNS
query. Deliver a canned response to the query and return whatever the
L{Deferred} returned by L{Resolver._query} fires with.
@param filter: The value to pass for the C{filter} parameter to
L{Resolver._query}.
"""
reactor = MemoryReactor()
resolver = Resolver([], reactor=reactor)
d = resolver._query(
Query(b'foo.example.com', A, IN), [('1.1.2.3', 1053)], (30,),
filter)
# A UDP port should have been started.
portNumber, transport = reactor.udpPorts.popitem()
# And a DNS packet sent.
[(packet, address)] = transport._sentPackets
message = Message()
message.fromStr(packet)
# It should be a query with the parameters used above.
self.assertEqual(message.queries, [Query(b'foo.example.com', A, IN)])
self.assertEqual(message.answers, [])
self.assertEqual(message.authority, [])
self.assertEqual(message.additional, [])
response = []
d.addCallback(response.append)
self.assertEqual(response, [])
# Once a reply is received, the Deferred should fire.
del message.queries[:]
message.answer = 1
message.answers.append(RRHeader(
b'foo.example.com', payload=Record_A('5.8.13.21')))
transport._protocol.datagramReceived(
message.toStr(), ('1.1.2.3', 1053))
return response[0]
def test_filteredQuery(self):
"""
L{Resolver._query} accepts a L{Query} instance and an address, issues
the query, and returns a L{Deferred} which fires with the response to
the query. If a true value is passed for the C{filter} parameter, the
result is a three-tuple of lists of records.
"""
answer, authority, additional = self._queryTest(True)
self.assertEqual(
answer,
[RRHeader(b'foo.example.com', payload=Record_A('5.8.13.21', ttl=0))])
self.assertEqual(authority, [])
self.assertEqual(additional, [])
def test_unfilteredQuery(self):
"""
Similar to L{test_filteredQuery}, but for the case where a false value
is passed for the C{filter} parameter. In this case, the result is a
L{Message} instance.
"""
message = self._queryTest(False)
self.assertIsInstance(message, Message)
self.assertEqual(message.queries, [])
self.assertEqual(
message.answers,
[RRHeader(b'foo.example.com', payload=Record_A('5.8.13.21', ttl=0))])
self.assertEqual(message.authority, [])
self.assertEqual(message.additional, [])
def _respond(self, answers=[], authority=[], additional=[], rCode=OK):
"""
Create a L{Message} suitable for use as a response to a query.
@param answers: A C{list} of two-tuples giving data for the answers
section of the message. The first element of each tuple is a name
for the L{RRHeader}. The second element is the payload.
@param authority: A C{list} like C{answers}, but for the authority
section of the response.
@param additional: A C{list} like C{answers}, but for the
additional section of the response.
@param rCode: The response code the message will be created with.
@return: A new L{Message} initialized with the given values.
"""
response = Message(rCode=rCode)
for (section, data) in [(response.answers, answers),
(response.authority, authority),
(response.additional, additional)]:
section.extend([
RRHeader(name, record.TYPE, getattr(record, 'CLASS', IN),
payload=record)
for (name, record) in data])
return response
def _getResolver(self, serverResponses, maximumQueries=10):
"""
Create and return a new L{root.Resolver} modified to resolve queries
against the record data represented by C{servers}.
@param serverResponses: A mapping from dns server addresses to
mappings. The inner mappings are from query two-tuples (name,
type) to dictionaries suitable for use as **arguments to
L{_respond}. See that method for details.
"""
roots = ['1.1.2.3']
resolver = Resolver(roots, maximumQueries)
def query(query, serverAddresses, timeout, filter):
msg("Query for QNAME %s at %r" % (query.name, serverAddresses))
for addr in serverAddresses:
try:
server = serverResponses[addr]
except KeyError:
continue
records = server[query.name.name, query.type]
return succeed(self._respond(**records))
resolver._query = query
return resolver
def test_lookupAddress(self):
"""
L{root.Resolver.lookupAddress} looks up the I{A} records for the
specified hostname by first querying one of the root servers the
resolver was created with and then following the authority delegations
until a result is received.
"""
servers = {
('1.1.2.3', 53): {
(b'foo.example.com', A): {
'authority': [(b'foo.example.com', Record_NS(b'ns1.example.com'))],
'additional': [(b'ns1.example.com', Record_A('34.55.89.144'))],
},
},
('34.55.89.144', 53): {
(b'foo.example.com', A): {
'answers': [(b'foo.example.com', Record_A('10.0.0.1'))],
}
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'foo.example.com')
d.addCallback(getOneAddress)
d.addCallback(self.assertEqual, '10.0.0.1')
return d
def test_lookupChecksClass(self):
"""
If a response includes a record with a class different from the one
in the query, it is ignored and lookup continues until a record with
the right class is found.
"""
badClass = Record_A('10.0.0.1')
badClass.CLASS = HS
servers = {
('1.1.2.3', 53): {
(b'foo.example.com', A): {
'answers': [(b'foo.example.com', badClass)],
'authority': [(b'foo.example.com', Record_NS(b'ns1.example.com'))],
'additional': [(b'ns1.example.com', Record_A('10.0.0.2'))],
},
},
('10.0.0.2', 53): {
(b'foo.example.com', A): {
'answers': [(b'foo.example.com', Record_A('10.0.0.3'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'foo.example.com')
d.addCallback(getOnePayload)
d.addCallback(self.assertEqual, Record_A('10.0.0.3'))
return d
def test_missingGlue(self):
"""
If an intermediate response includes no glue records for the
authorities, separate queries are made to find those addresses.
"""
servers = {
('1.1.2.3', 53): {
(b'foo.example.com', A): {
'authority': [(b'foo.example.com', Record_NS(b'ns1.example.org'))],
# Conspicuous lack of an additional section naming ns1.example.com
},
(b'ns1.example.org', A): {
'answers': [(b'ns1.example.org', Record_A('10.0.0.1'))],
},
},
('10.0.0.1', 53): {
(b'foo.example.com', A): {
'answers': [(b'foo.example.com', Record_A('10.0.0.2'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'foo.example.com')
d.addCallback(getOneAddress)
d.addCallback(self.assertEqual, '10.0.0.2')
return d
def test_missingName(self):
"""
If a name is missing, L{Resolver.lookupAddress} returns a L{Deferred}
which fails with L{DNSNameError}.
"""
servers = {
('1.1.2.3', 53): {
(b'foo.example.com', A): {
'rCode': ENAME,
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'foo.example.com')
return self.assertFailure(d, DNSNameError)
def test_answerless(self):
"""
If a query is responded to with no answers or nameserver records, the
L{Deferred} returned by L{Resolver.lookupAddress} fires with
L{ResolverError}.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
return self.assertFailure(d, ResolverError)
def test_delegationLookupError(self):
"""
If there is an error resolving the nameserver in a delegation response,
the L{Deferred} returned by L{Resolver.lookupAddress} fires with that
error.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'authority': [(b'example.com', Record_NS(b'ns1.example.com'))],
},
(b'ns1.example.com', A): {
'rCode': ENAME,
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
return self.assertFailure(d, DNSNameError)
def test_delegationLookupEmpty(self):
"""
If there are no records in the response to a lookup of a delegation
nameserver, the L{Deferred} returned by L{Resolver.lookupAddress} fires
with L{ResolverError}.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'authority': [(b'example.com', Record_NS(b'ns1.example.com'))],
},
(b'ns1.example.com', A): {
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
return self.assertFailure(d, ResolverError)
def test_lookupNameservers(self):
"""
L{Resolver.lookupNameservers} is like L{Resolver.lookupAddress}, except
it queries for I{NS} records instead of I{A} records.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'rCode': ENAME,
},
(b'example.com', NS): {
'answers': [(b'example.com', Record_NS(b'ns1.example.com'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupNameservers(b'example.com')
def getOneName(results):
ans, auth, add = results
return ans[0].payload.name
d.addCallback(getOneName)
d.addCallback(self.assertEqual, Name(b'ns1.example.com'))
return d
def test_returnCanonicalName(self):
"""
If a I{CNAME} record is encountered as the answer to a query for
another record type, that record is returned as the answer.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'answers': [(b'example.com', Record_CNAME(b'example.net')),
(b'example.net', Record_A('10.0.0.7'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
d.addCallback(lambda results: results[0]) # Get the answer section
d.addCallback(
self.assertEqual,
[RRHeader(b'example.com', CNAME, payload=Record_CNAME(b'example.net')),
RRHeader(b'example.net', A, payload=Record_A('10.0.0.7'))])
return d
def test_followCanonicalName(self):
"""
If no record of the requested type is included in a response, but a
I{CNAME} record for the query name is included, queries are made to
resolve the value of the I{CNAME}.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'answers': [(b'example.com', Record_CNAME(b'example.net'))],
},
(b'example.net', A): {
'answers': [(b'example.net', Record_A('10.0.0.5'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
d.addCallback(lambda results: results[0]) # Get the answer section
d.addCallback(
self.assertEqual,
[RRHeader(b'example.com', CNAME, payload=Record_CNAME(b'example.net')),
RRHeader(b'example.net', A, payload=Record_A('10.0.0.5'))])
return d
def test_detectCanonicalNameLoop(self):
"""
If there is a cycle between I{CNAME} records in a response, this is
detected and the L{Deferred} returned by the lookup method fails
with L{ResolverError}.
"""
servers = {
('1.1.2.3', 53): {
(b'example.com', A): {
'answers': [(b'example.com', Record_CNAME(b'example.net')),
(b'example.net', Record_CNAME(b'example.com'))],
},
},
}
resolver = self._getResolver(servers)
d = resolver.lookupAddress(b'example.com')
return self.assertFailure(d, ResolverError)
def test_boundedQueries(self):
"""
L{Resolver.lookupAddress} won't issue more queries following
delegations than the limit passed to its initializer.
"""
servers = {
('1.1.2.3', 53): {
# First query - force it to start over with a name lookup of
# ns1.example.com
(b'example.com', A): {
'authority': [(b'example.com', Record_NS(b'ns1.example.com'))],
},
# Second query - let it resume the original lookup with the
# address of the nameserver handling the delegation.
(b'ns1.example.com', A): {
'answers': [(b'ns1.example.com', Record_A('10.0.0.2'))],
},
},
('10.0.0.2', 53): {
# Third query - let it jump straight to asking the
# delegation server by including its address here (different
# case from the first query).
(b'example.com', A): {
'authority': [(b'example.com', Record_NS(b'ns2.example.com'))],
'additional': [(b'ns2.example.com', Record_A('10.0.0.3'))],
},
},
('10.0.0.3', 53): {
# Fourth query - give it the answer, we're done.
(b'example.com', A): {
'answers': [(b'example.com', Record_A('10.0.0.4'))],
},
},
}
# Make two resolvers. One which is allowed to make 3 queries
# maximum, and so will fail, and on which may make 4, and so should
# succeed.
failer = self._getResolver(servers, 3)
failD = self.assertFailure(
failer.lookupAddress(b'example.com'), ResolverError)
succeeder = self._getResolver(servers, 4)
succeedD = succeeder.lookupAddress(b'example.com')
succeedD.addCallback(getOnePayload)
succeedD.addCallback(self.assertEqual, Record_A('10.0.0.4'))
return gatherResults([failD, succeedD])
class ResolverFactoryArguments(Exception):
"""
Raised by L{raisingResolverFactory} with the *args and **kwargs passed to
that function.
"""
def __init__(self, args, kwargs):
"""
Store the supplied args and kwargs as attributes.
@param args: Positional arguments.
@param kwargs: Keyword arguments.
"""
self.args = args
self.kwargs = kwargs
def raisingResolverFactory(*args, **kwargs):
"""
Raise a L{ResolverFactoryArguments} exception containing the
positional and keyword arguments passed to resolverFactory.
@param args: A L{list} of all the positional arguments supplied by
the caller.
@param kwargs: A L{list} of all the keyword arguments supplied by
the caller.
"""
raise ResolverFactoryArguments(args, kwargs)
class RootResolverResolverFactoryTests(TestCase):
"""
Tests for L{root.Resolver._resolverFactory}.
"""
def test_resolverFactoryArgumentPresent(self):
"""
L{root.Resolver.__init__} accepts a C{resolverFactory}
argument and assigns it to C{self._resolverFactory}.
"""
r = Resolver(hints=[None], resolverFactory=raisingResolverFactory)
self.assertIs(r._resolverFactory, raisingResolverFactory)
def test_resolverFactoryArgumentAbsent(self):
"""
L{root.Resolver.__init__} sets L{client.Resolver} as the
C{_resolverFactory} if a C{resolverFactory} argument is not
supplied.
"""
r = Resolver(hints=[None])
self.assertIs(r._resolverFactory, client.Resolver)
def test_resolverFactoryOnlyExpectedArguments(self):
"""
L{root.Resolver._resolverFactory} is supplied with C{reactor} and
C{servers} keyword arguments.
"""
dummyReactor = object()
r = Resolver(hints=['192.0.2.101'],
resolverFactory=raisingResolverFactory,
reactor=dummyReactor)
e = self.assertRaises(ResolverFactoryArguments,
r.lookupAddress, 'example.com')
self.assertEqual(
((), {'reactor': dummyReactor, 'servers': [('192.0.2.101', 53)]}),
(e.args, e.kwargs)
)
ROOT_SERVERS = [
'a.root-servers.net',
'b.root-servers.net',
'c.root-servers.net',
'd.root-servers.net',
'e.root-servers.net',
'f.root-servers.net',
'g.root-servers.net',
'h.root-servers.net',
'i.root-servers.net',
'j.root-servers.net',
'k.root-servers.net',
'l.root-servers.net',
'm.root-servers.net']
@implementer(IResolverSimple)
class StubResolver(object):
"""
An L{IResolverSimple} implementer which traces all getHostByName
calls and their deferred results. The deferred results can be
accessed and fired synchronously.
"""
def __init__(self):
"""
@type calls: L{list} of L{tuple} containing C{args} and
C{kwargs} supplied to C{getHostByName} calls.
@type pendingResults: L{list} of L{Deferred} returned by
C{getHostByName}.
"""
self.calls = []
self.pendingResults = []
def getHostByName(self, *args, **kwargs):
"""
A fake implementation of L{IResolverSimple.getHostByName}
@param args: A L{list} of all the positional arguments supplied by
the caller.
@param kwargs: A L{list} of all the keyword arguments supplied by
the caller.
@return: A L{Deferred} which may be fired later from the test
fixture.
"""
self.calls.append((args, kwargs))
d = Deferred()
self.pendingResults.append(d)
return d
verifyClass(IResolverSimple, StubResolver)
class BootstrapTests(SynchronousTestCase):
"""
Tests for L{root.bootstrap}
"""
def test_returnsDeferredResolver(self):
"""
L{root.bootstrap} returns an object which is initially a
L{root.DeferredResolver}.
"""
deferredResolver = root.bootstrap(StubResolver())
self.assertIsInstance(deferredResolver, root.DeferredResolver)
def test_resolves13RootServers(self):
"""
The L{IResolverSimple} supplied to L{root.bootstrap} is used to lookup
the IP addresses of the 13 root name servers.
"""
stubResolver = StubResolver()
root.bootstrap(stubResolver)
self.assertEqual(
stubResolver.calls,
[((s,), {}) for s in ROOT_SERVERS])
def test_becomesResolver(self):
"""
The L{root.DeferredResolver} initially returned by L{root.bootstrap}
becomes a L{root.Resolver} when the supplied resolver has successfully
looked up all root hints.
"""
stubResolver = StubResolver()
deferredResolver = root.bootstrap(stubResolver)
for d in stubResolver.pendingResults:
d.callback('192.0.2.101')
self.assertIsInstance(deferredResolver, Resolver)
def test_resolverReceivesRootHints(self):
"""
The L{root.Resolver} which eventually replaces L{root.DeferredResolver}
is supplied with the IP addresses of the 13 root servers.
"""
stubResolver = StubResolver()
deferredResolver = root.bootstrap(stubResolver)
for d in stubResolver.pendingResults:
d.callback('192.0.2.101')
self.assertEqual(deferredResolver.hints, ['192.0.2.101'] * 13)
def test_continuesWhenSomeRootHintsFail(self):
"""
The L{root.Resolver} is eventually created, even if some of the root
hint lookups fail. Only the working root hint IP addresses are supplied
to the L{root.Resolver}.
"""
stubResolver = StubResolver()
deferredResolver = root.bootstrap(stubResolver)
results = iter(stubResolver.pendingResults)
d1 = next(results)
for d in results:
d.callback('192.0.2.101')
d1.errback(TimeoutError())
def checkHints(res):
self.assertEqual(deferredResolver.hints, ['192.0.2.101'] * 12)
d1.addBoth(checkHints)
def test_continuesWhenAllRootHintsFail(self):
"""
The L{root.Resolver} is eventually created, even if all of the root hint
lookups fail. Pending and new lookups will then fail with
AttributeError.
"""
stubResolver = StubResolver()
deferredResolver = root.bootstrap(stubResolver)
results = iter(stubResolver.pendingResults)
d1 = next(results)
for d in results:
d.errback(TimeoutError())
d1.errback(TimeoutError())
def checkHints(res):
self.assertEqual(deferredResolver.hints, [])
d1.addBoth(checkHints)
self.addCleanup(self.flushLoggedErrors, TimeoutError)
def test_passesResolverFactory(self):
"""
L{root.bootstrap} accepts a C{resolverFactory} argument which is passed
as an argument to L{root.Resolver} when it has successfully looked up
root hints.
"""
stubResolver = StubResolver()
deferredResolver = root.bootstrap(
stubResolver, resolverFactory=raisingResolverFactory)
for d in stubResolver.pendingResults:
d.callback('192.0.2.101')
self.assertIs(
deferredResolver._resolverFactory, raisingResolverFactory)
class StubDNSDatagramProtocol:
"""
A do-nothing stand-in for L{DNSDatagramProtocol} which can be used to avoid
network traffic in tests where that kind of thing doesn't matter.
"""
def query(self, *a, **kw):
return Deferred()
_retrySuppression = util.suppress(
category=DeprecationWarning,
message=(
'twisted.names.root.retry is deprecated since Twisted 10.0. Use a '
'Resolver object for retry logic.'))
@@ -0,0 +1,137 @@
# Copyright (c) Twisted Matrix Laboratories.
# See LICENSE for details.
"""
Utilities for Twisted.names tests.
"""
from __future__ import division, absolute_import
from random import randrange
from zope.interface import implementer
from zope.interface.verify import verifyClass
from twisted.internet.address import IPv4Address
from twisted.internet.defer import succeed
from twisted.internet.task import Clock
from twisted.internet.interfaces import IReactorUDP, IUDPTransport
@implementer(IUDPTransport)
class MemoryDatagramTransport(object):
"""
This L{IUDPTransport} implementation enforces the usual connection rules
and captures sent traffic in a list for later inspection.
@ivar _host: The host address to which this transport is bound.
@ivar _protocol: The protocol connected to this transport.
@ivar _sentPackets: A C{list} of two-tuples of the datagrams passed to
C{write} and the addresses to which they are destined.
@ivar _connectedTo: L{None} if this transport is unconnected, otherwise an
address to which all traffic is supposedly sent.
@ivar _maxPacketSize: An C{int} giving the maximum length of a datagram
which will be successfully handled by C{write}.
"""
def __init__(self, host, protocol, maxPacketSize):
self._host = host
self._protocol = protocol
self._sentPackets = []
self._connectedTo = None
self._maxPacketSize = maxPacketSize
def getHost(self):
"""
Return the address which this transport is pretending to be bound
to.
"""
return IPv4Address('UDP', *self._host)
def connect(self, host, port):
"""
Connect this transport to the given address.
"""
if self._connectedTo is not None:
raise ValueError("Already connected")
self._connectedTo = (host, port)
def write(self, datagram, addr=None):
"""
Send the given datagram.
"""
if addr is None:
addr = self._connectedTo
if addr is None:
raise ValueError("Need an address")
if len(datagram) > self._maxPacketSize:
raise ValueError("Packet too big")
self._sentPackets.append((datagram, addr))
def stopListening(self):
"""
Shut down this transport.
"""
self._protocol.stopProtocol()
return succeed(None)
def setBroadcastAllowed(self, enabled):
"""
Dummy implementation to satisfy L{IUDPTransport}.
"""
pass
def getBroadcastAllowed(self):
"""
Dummy implementation to satisfy L{IUDPTransport}.
"""
pass
verifyClass(IUDPTransport, MemoryDatagramTransport)
@implementer(IReactorUDP)
class MemoryReactor(Clock):
"""
An L{IReactorTime} and L{IReactorUDP} provider.
Time is controlled deterministically via the base class, L{Clock}. UDP is
handled in-memory by connecting protocols to instances of
L{MemoryDatagramTransport}.
@ivar udpPorts: A C{dict} mapping port numbers to instances of
L{MemoryDatagramTransport}.
"""
def __init__(self):
Clock.__init__(self)
self.udpPorts = {}
def listenUDP(self, port, protocol, interface='', maxPacketSize=8192):
"""
Pretend to bind a UDP port and connect the given protocol to it.
"""
if port == 0:
while True:
port = randrange(1, 2 ** 16)
if port not in self.udpPorts:
break
if port in self.udpPorts:
raise ValueError("Address in use")
transport = MemoryDatagramTransport(
(interface, port), protocol, maxPacketSize)
self.udpPorts[port] = transport
protocol.makeConnection(transport)
return transport
verifyClass(IReactorUDP, MemoryReactor)