mirror of https://github.com/scrapy/scrapy.git
Fix port issue with cached DNS (#7772)
* Fix port issue with cached DNS * Keep ports off the cache * Complete coverage
This commit is contained in:
parent
25e6884e2f
commit
f3693aa8ba
|
|
@ -2,6 +2,7 @@ from __future__ import annotations
|
|||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import attr
|
||||
from twisted.internet import defer
|
||||
from twisted.internet.base import ReactorBase, ThreadedResolver
|
||||
from twisted.internet.interfaces import (
|
||||
|
|
@ -70,6 +71,12 @@ class CachingThreadedResolver(ThreadedResolver):
|
|||
return result
|
||||
|
||||
|
||||
def _address_with_port(address: IAddress, port: int) -> IAddress:
|
||||
if getattr(address, "port", port) == port:
|
||||
return address
|
||||
return attr.evolve(address, port=port)
|
||||
|
||||
|
||||
@implementer(IHostResolution)
|
||||
class HostResolution:
|
||||
def __init__(self, name: str):
|
||||
|
|
@ -97,7 +104,11 @@ class _CachingResolutionReceiver:
|
|||
def resolutionComplete(self) -> None:
|
||||
self.resolutionReceiver.resolutionComplete()
|
||||
if self.addresses:
|
||||
dnscache[self.hostName] = self.addresses
|
||||
# Name resolution does not depend on the port, so cache entries are
|
||||
# kept port-agnostic and the requested port is set on cache hits.
|
||||
dnscache[self.hostName] = [
|
||||
_address_with_port(address, 0) for address in self.addresses
|
||||
]
|
||||
|
||||
|
||||
@implementer(IHostnameResolver)
|
||||
|
|
@ -142,7 +153,7 @@ class CachingHostnameResolver:
|
|||
transportSemantics,
|
||||
)
|
||||
resolutionReceiver.resolutionBegan(HostResolution(hostName))
|
||||
for addr in addresses:
|
||||
resolutionReceiver.addressResolved(addr)
|
||||
for address in addresses:
|
||||
resolutionReceiver.addressResolved(_address_with_port(address, portNumber))
|
||||
resolutionReceiver.resolutionComplete()
|
||||
return resolutionReceiver
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from __future__ import annotations
|
|||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from twisted.internet.address import IPv4Address, IPv6Address
|
||||
|
||||
from scrapy.resolver import CachingHostnameResolver, CachingThreadedResolver, dnscache
|
||||
from scrapy.utils.defer import maybe_deferred_to_future
|
||||
|
|
@ -55,11 +56,77 @@ def test_caching_hostname_resolver_no_addresses_not_cached():
|
|||
assert "example.com" not in dnscache
|
||||
|
||||
|
||||
def test_caching_hostname_resolver_cached_addresses_have_no_port():
|
||||
def fake_resolve(receiver, *_):
|
||||
receiver.resolutionBegan(Mock())
|
||||
receiver.addressResolved(IPv4Address("TCP", "1.2.3.4", 80))
|
||||
receiver.addressResolved(IPv6Address("TCP", "::1", 80))
|
||||
receiver.resolutionComplete()
|
||||
return receiver
|
||||
|
||||
reactor = Mock()
|
||||
reactor.nameResolver.resolveHostName.side_effect = fake_resolve
|
||||
|
||||
receiver = Mock()
|
||||
resolver = CachingHostnameResolver(reactor, cache_size=10)
|
||||
resolver.resolveHostName(receiver, "example.com", portNumber=80)
|
||||
|
||||
# The port requested on a cache miss is passed through unchanged, but it is
|
||||
# not part of what gets cached.
|
||||
resolved_ports = [
|
||||
call.args[0].port for call in receiver.addressResolved.call_args_list
|
||||
]
|
||||
assert resolved_ports == [80, 80]
|
||||
assert [address.port for address in dnscache["example.com"]] == [0, 0]
|
||||
|
||||
|
||||
def test_caching_hostname_resolver_cache_hit_without_port():
|
||||
cached_addresses = [
|
||||
IPv4Address("TCP", "1.2.3.4", 0),
|
||||
IPv6Address("TCP", "::1", 0),
|
||||
]
|
||||
dnscache["example.com"] = cached_addresses
|
||||
|
||||
receiver = Mock()
|
||||
resolver = CachingHostnameResolver(Mock(), cache_size=10)
|
||||
resolver.resolveHostName(receiver, "example.com")
|
||||
|
||||
# Cached addresses already use the requested port, so they are reused as is.
|
||||
resolved_addresses = [
|
||||
call.args[0] for call in receiver.addressResolved.call_args_list
|
||||
]
|
||||
assert all(
|
||||
resolved is cached
|
||||
for resolved, cached in zip(resolved_addresses, cached_addresses, strict=True)
|
||||
)
|
||||
|
||||
|
||||
def test_caching_hostname_resolver_cache_hit_uses_requested_port():
|
||||
dnscache["example.com"] = [
|
||||
IPv4Address("TCP", "1.2.3.4", 0),
|
||||
IPv6Address("TCP", "::1", 0),
|
||||
]
|
||||
|
||||
receiver = Mock()
|
||||
resolver = CachingHostnameResolver(Mock(), cache_size=10)
|
||||
resolver.resolveHostName(receiver, "example.com", portNumber=443)
|
||||
|
||||
resolved_addresses = [
|
||||
call.args[0] for call in receiver.addressResolved.call_args_list
|
||||
]
|
||||
assert resolved_addresses == [
|
||||
IPv4Address("TCP", "1.2.3.4", 443),
|
||||
IPv6Address("TCP", "::1", 443),
|
||||
]
|
||||
# The cached addresses must not be mutated in place.
|
||||
assert [address.port for address in dnscache["example.com"]] == [0, 0]
|
||||
|
||||
|
||||
def test_caching_hostname_resolver_dnscache_disabled_rejects_storage():
|
||||
|
||||
def fake_resolve(receiver, *_):
|
||||
receiver.resolutionBegan(Mock())
|
||||
receiver.addressResolved(Mock())
|
||||
receiver.addressResolved(IPv4Address("TCP", "1.2.3.4", 80))
|
||||
receiver.resolutionComplete()
|
||||
return receiver
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue