Fix CachingHostnameResolver

This commit is contained in:
Eugenio Lacuesta 2020-09-21 10:59:55 -03:00
parent da426fb3cf
commit b55c911ddc
No known key found for this signature in database
GPG Key ID: DA3EF2D0913E9810
1 changed files with 34 additions and 23 deletions

View File

@ -1,4 +1,5 @@
from twisted.internet import defer
from twisted.internet._resolver import HostResolution
from twisted.internet.base import ThreadedResolver
from twisted.internet.interfaces import IHostnameResolver, IResolutionReceiver, IResolverSimple
from zope.interface.declarations import implementer, provider
@ -50,6 +51,27 @@ class CachingThreadedResolver(ThreadedResolver):
return result
@provider(IResolutionReceiver)
class _CachingResolutionReceiver:
def __init__(self, resolutionReceiver, hostName):
self.resolutionReceiver = resolutionReceiver
self.hostName = hostName
self.addresses = []
def resolutionBegan(self, resolution):
self.resolutionReceiver.resolutionBegan(resolution)
self.resolution = resolution
def addressResolved(self, address):
self.resolutionReceiver.addressResolved(address)
self.addresses.append(address)
def resolutionComplete(self):
self.resolutionReceiver.resolutionComplete()
if self.addresses:
dnscache[self.hostName] = self.addresses
@implementer(IHostnameResolver)
class CachingHostnameResolver:
"""
@ -73,33 +95,22 @@ class CachingHostnameResolver:
def install_on_reactor(self):
self.reactor.installNameResolver(self)
def resolveHostName(self, resolutionReceiver, hostName, portNumber=0,
addressTypes=None, transportSemantics='TCP'):
@provider(IResolutionReceiver)
class CachingResolutionReceiver(resolutionReceiver):
def resolutionBegan(self, resolution):
super().resolutionBegan(resolution)
self.resolution = resolution
self.resolved = False
def addressResolved(self, address):
super().addressResolved(address)
self.resolved = True
def resolutionComplete(self):
super().resolutionComplete()
if self.resolved:
dnscache[hostName] = self.resolution
def resolveHostName(
self, resolutionReceiver, hostName, portNumber=0, addressTypes=None, transportSemantics="TCP"
):
try:
return dnscache[hostName]
addresses = dnscache[hostName]
except KeyError:
return self.original_resolver.resolveHostName(
CachingResolutionReceiver(),
_CachingResolutionReceiver(resolutionReceiver, hostName),
hostName,
portNumber,
addressTypes,
transportSemantics
transportSemantics,
)
else:
resolutionReceiver.resolutionBegan(HostResolution(hostName))
for addr in addresses:
resolutionReceiver.addressResolved(addr)
resolutionReceiver.resolutionComplete()
return resolutionReceiver