17.12
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user