diff --git a/docs/topics/asyncio.rst b/docs/topics/asyncio.rst index 91e1cca0d..82c5f271f 100644 --- a/docs/topics/asyncio.rst +++ b/docs/topics/asyncio.rst @@ -6,13 +6,14 @@ asyncio .. versionadded:: 2.0 -Scrapy has partial support :mod:`asyncio`. After you :ref:`install the asyncio -reactor `, you may use :mod:`asyncio` and +Scrapy has partial support for :mod:`asyncio`. After you :ref:`install the +asyncio reactor `, you may use :mod:`asyncio` and :mod:`asyncio`-powered libraries in any :doc:`coroutine `. -.. warning:: :mod:`asyncio` support in Scrapy is experimental. Future Scrapy - versions may introduce related changes without a deprecation - period or warning. +.. warning:: :mod:`asyncio` support in Scrapy is experimental, and not yet + recommended for production environments. Future Scrapy versions + may introduce related changes without a deprecation period or + warning. .. _install-asyncio: diff --git a/docs/topics/commands.rst b/docs/topics/commands.rst index 7de5e8121..eef6b36ff 100644 --- a/docs/topics/commands.rst +++ b/docs/topics/commands.rst @@ -598,8 +598,6 @@ Example: Register commands via setup.py entry points ------------------------------------------- -.. note:: This is an experimental feature, use with caution. - You can also add Scrapy commands from an external library by adding a ``scrapy.commands`` section in the entry points of the library ``setup.py`` file. diff --git a/docs/topics/request-response.rst b/docs/topics/request-response.rst index 37008f3e9..a0448c5ab 100644 --- a/docs/topics/request-response.rst +++ b/docs/topics/request-response.rst @@ -694,7 +694,7 @@ Response objects :type ip_address: :class:`ipaddress.IPv4Address` or :class:`ipaddress.IPv6Address` :param protocol: The protocol that was used to download the response. - For instance: "HTTP/1.0", "HTTP/1.1" + For instance: "HTTP/1.0", "HTTP/1.1", "h2" :type protocol: :class:`str` .. versionadded:: 2.0.0 diff --git a/docs/topics/settings.rst b/docs/topics/settings.rst index 0086a6c74..0a4684a91 100644 --- a/docs/topics/settings.rst +++ b/docs/topics/settings.rst @@ -677,6 +677,38 @@ handler (without replacement), place this in your ``settings.py``:: 'ftp': None, } +The default HTTPS handler uses HTTP/1.1. To use HTTP/2 update +:setting:`DOWNLOAD_HANDLERS` as follows:: + + DOWNLOAD_HANDLERS = { + 'https': 'scrapy.core.downloader.handlers.http2.H2DownloadHandler', + } + +.. warning:: + + HTTP/2 support in Scrapy is experimental, and not yet recommended for + production environments. Future Scrapy versions may introduce related + changes without a deprecation period or warning. + +.. note:: + + Known limitations of the current HTTP/2 implementation of Scrapy include: + + - No support for HTTP/2 Cleartext (h2c), since no major browser supports + HTTP/2 unencrypted (refer `http2 faq`_). + + - No setting to specify a maximum `frame size`_ larger than the default + value, 16384. Connections to servers that send a larger frame will + fail. + + - No support for `server pushes`_, which are ignored. + + - No support for the :signal:`bytes_received` signal. + +.. _frame size: https://tools.ietf.org/html/rfc7540#section-4.2 +.. _http2 faq: https://http2.github.io/faq/#does-http2-require-encryption +.. _server pushes: https://tools.ietf.org/html/rfc7540#section-8.2 + .. setting:: DOWNLOAD_TIMEOUT DOWNLOAD_TIMEOUT @@ -754,6 +786,15 @@ Optionally, this can be set per-request basis by using the If :setting:`RETRY_ENABLED` is ``True`` and this setting is set to ``True``, the ``ResponseFailed([_DataLoss])`` failure will be retried as usual. +.. warning:: + + This setting is ignored by the + :class:`~scrapy.core.downloader.handlers.http2.H2DownloadHandler` + download handler (see :setting:`DOWNLOAD_HANDLERS`). In case of a data loss + error, the corresponding HTTP/2 connection may be corrupted, affecting other + requests that use the same connection; hence, a ``ResponseFailed([InvalidBodyLengthError])`` + failure is always raised for every request that was using that connection. + .. setting:: DUPEFILTER_CLASS DUPEFILTER_CLASS diff --git a/scrapy/core/downloader/contextfactory.py b/scrapy/core/downloader/contextfactory.py index 8a7d656a1..073ef16bf 100644 --- a/scrapy/core/downloader/contextfactory.py +++ b/scrapy/core/downloader/contextfactory.py @@ -1,10 +1,15 @@ +import warnings + from OpenSSL import SSL +from twisted.internet._sslverify import _setAcceptableProtocols from twisted.internet.ssl import optionsForClientTLS, CertificateOptions, platformTrust, AcceptableCiphers from twisted.web.client import BrowserLikePolicyForHTTPS from twisted.web.iweb import IPolicyForHTTPS from zope.interface.declarations import implementer +from zope.interface.verify import verifyObject -from scrapy.core.downloader.tls import ScrapyClientTLSOptions, DEFAULT_CIPHERS +from scrapy.core.downloader.tls import DEFAULT_CIPHERS, openssl_methods, ScrapyClientTLSOptions +from scrapy.utils.misc import create_instance, load_object @implementer(IPolicyForHTTPS) @@ -81,8 +86,8 @@ class BrowserLikeContextFactory(ScrapyClientContextFactory): The default OpenSSL method is ``TLS_METHOD`` (also called ``SSLv23_METHOD``) which allows TLS protocol negotiation. """ - def creatorForNetloc(self, hostname, port): + def creatorForNetloc(self, hostname, port): # trustRoot set to platformTrust() will use the platform's root CAs. # # This means that a website like https://www.cacert.org will be rejected @@ -92,3 +97,49 @@ class BrowserLikeContextFactory(ScrapyClientContextFactory): trustRoot=platformTrust(), extraCertificateOptions={'method': self._ssl_method}, ) + + +@implementer(IPolicyForHTTPS) +class AcceptableProtocolsContextFactory: + """Context factory to used to override the acceptable protocols + to set up the [OpenSSL.SSL.Context] for doing NPN and/or ALPN + negotiation. + """ + + def __init__(self, context_factory, acceptable_protocols): + verifyObject(IPolicyForHTTPS, context_factory) + self._wrapped_context_factory = context_factory + self._acceptable_protocols = acceptable_protocols + + def creatorForNetloc(self, hostname, port): + options = self._wrapped_context_factory.creatorForNetloc(hostname, port) + _setAcceptableProtocols(options._ctx, self._acceptable_protocols) + return options + + +def load_context_factory_from_settings(settings, crawler): + ssl_method = openssl_methods[settings.get('DOWNLOADER_CLIENT_TLS_METHOD')] + context_factory_cls = load_object(settings['DOWNLOADER_CLIENTCONTEXTFACTORY']) + # try method-aware context factory + try: + context_factory = create_instance( + objcls=context_factory_cls, + settings=settings, + crawler=crawler, + method=ssl_method, + ) + except TypeError: + # use context factory defaults + context_factory = create_instance( + objcls=context_factory_cls, + settings=settings, + crawler=crawler, + ) + msg = """ + '%s' does not accept `method` argument (type OpenSSL.SSL method,\ + e.g. OpenSSL.SSL.SSLv23_METHOD) and/or `tls_verbose_logging` argument and/or `tls_ciphers` argument.\ + Please upgrade your context factory class to handle them or ignore them.""" % ( + settings['DOWNLOADER_CLIENTCONTEXTFACTORY'],) + warnings.warn(msg) + + return context_factory diff --git a/scrapy/core/downloader/handlers/http11.py b/scrapy/core/downloader/handlers/http11.py index 516a4326b..9cdadb27f 100644 --- a/scrapy/core/downloader/handlers/http11.py +++ b/scrapy/core/downloader/handlers/http11.py @@ -20,12 +20,11 @@ from twisted.web.iweb import IBodyProducer, UNKNOWN_LENGTH from zope.interface import implementer from scrapy import signals -from scrapy.core.downloader.tls import openssl_methods +from scrapy.core.downloader.contextfactory import load_context_factory_from_settings from scrapy.core.downloader.webclient import _parse from scrapy.exceptions import ScrapyDeprecationWarning, StopDownload from scrapy.http import Headers from scrapy.responsetypes import responsetypes -from scrapy.utils.misc import create_instance, load_object from scrapy.utils.python import to_bytes, to_unicode @@ -43,29 +42,7 @@ class HTTP11DownloadHandler: self._pool.maxPersistentPerHost = settings.getint('CONCURRENT_REQUESTS_PER_DOMAIN') self._pool._factory.noisy = False - self._sslMethod = openssl_methods[settings.get('DOWNLOADER_CLIENT_TLS_METHOD')] - self._contextFactoryClass = load_object(settings['DOWNLOADER_CLIENTCONTEXTFACTORY']) - # try method-aware context factory - try: - self._contextFactory = create_instance( - objcls=self._contextFactoryClass, - settings=settings, - crawler=crawler, - method=self._sslMethod, - ) - except TypeError: - # use context factory defaults - self._contextFactory = create_instance( - objcls=self._contextFactoryClass, - settings=settings, - crawler=crawler, - ) - msg = f""" - '{settings["DOWNLOADER_CLIENTCONTEXTFACTORY"]}' does not accept `method` \ - argument (type OpenSSL.SSL method, e.g. OpenSSL.SSL.SSLv23_METHOD) and/or \ - `tls_verbose_logging` argument and/or `tls_ciphers` argument.\ - Please upgrade your context factory class to handle them or ignore them.""" - warnings.warn(msg) + self._contextFactory = load_context_factory_from_settings(settings, crawler) self._default_maxsize = settings.getint('DOWNLOAD_MAXSIZE') self._default_warnsize = settings.getint('DOWNLOAD_WARNSIZE') self._fail_on_dataloss = settings.getbool('DOWNLOAD_FAIL_ON_DATALOSS') diff --git a/scrapy/core/downloader/handlers/http2.py b/scrapy/core/downloader/handlers/http2.py new file mode 100644 index 000000000..e97c31e90 --- /dev/null +++ b/scrapy/core/downloader/handlers/http2.py @@ -0,0 +1,125 @@ +import warnings +from time import time +from typing import Optional, Type, TypeVar +from urllib.parse import urldefrag + +from twisted.internet.base import DelayedCall +from twisted.internet.defer import Deferred +from twisted.internet.error import TimeoutError +from twisted.web.client import URI + +from scrapy.core.downloader.contextfactory import load_context_factory_from_settings +from scrapy.core.downloader.webclient import _parse +from scrapy.core.http2.agent import H2Agent, H2ConnectionPool, ScrapyProxyH2Agent +from scrapy.crawler import Crawler +from scrapy.http import Request, Response +from scrapy.settings import Settings +from scrapy.spiders import Spider +from scrapy.utils.python import to_bytes + + +H2DownloadHandlerOrSubclass = TypeVar("H2DownloadHandlerOrSubclass", bound="H2DownloadHandler") + + +class H2DownloadHandler: + def __init__(self, settings: Settings, crawler: Optional[Crawler] = None): + self._crawler = crawler + + from twisted.internet import reactor + self._pool = H2ConnectionPool(reactor, settings) + self._context_factory = load_context_factory_from_settings(settings, crawler) + + @classmethod + def from_crawler(cls: Type[H2DownloadHandlerOrSubclass], crawler: Crawler) -> H2DownloadHandlerOrSubclass: + return cls(crawler.settings, crawler) + + def download_request(self, request: Request, spider: Spider) -> Deferred: + agent = ScrapyH2Agent( + context_factory=self._context_factory, + pool=self._pool, + crawler=self._crawler, + ) + return agent.download_request(request, spider) + + def close(self) -> None: + self._pool.close_connections() + + +class ScrapyH2Agent: + _Agent = H2Agent + _ProxyAgent = ScrapyProxyH2Agent + + def __init__( + self, context_factory, + pool: H2ConnectionPool, + connect_timeout: int = 10, + bind_address: Optional[bytes] = None, + crawler: Optional[Crawler] = None, + ) -> None: + self._context_factory = context_factory + self._connect_timeout = connect_timeout + self._bind_address = bind_address + self._pool = pool + self._crawler = crawler + + def _get_agent(self, request: Request, timeout: Optional[float]) -> H2Agent: + from twisted.internet import reactor + bind_address = request.meta.get('bindaddress') or self._bind_address + proxy = request.meta.get('proxy') + if proxy: + _, _, proxy_host, proxy_port, proxy_params = _parse(proxy) + scheme = _parse(request.url)[0] + proxy_host = proxy_host.decode() + omit_connect_tunnel = b'noconnect' in proxy_params + if omit_connect_tunnel: + warnings.warn("Using HTTPS proxies in the noconnect mode is not supported by the " + "downloader handler. If you use Crawlera, it doesn't require this " + "mode anymore, so you should update scrapy-crawlera to 1.3.0+ " + "and remove '?noconnect' from the Crawlera URL.") + + if scheme == b'https' and not omit_connect_tunnel: + # ToDo + raise NotImplementedError('Tunneling via CONNECT method using HTTP/2.0 is not yet supported') + return self._ProxyAgent( + reactor=reactor, + context_factory=self._context_factory, + proxy_uri=URI.fromBytes(to_bytes(proxy, encoding='ascii')), + connect_timeout=timeout, + bind_address=bind_address, + pool=self._pool, + ) + + return self._Agent( + reactor=reactor, + context_factory=self._context_factory, + connect_timeout=timeout, + bind_address=bind_address, + pool=self._pool, + ) + + def download_request(self, request: Request, spider: Spider) -> Deferred: + from twisted.internet import reactor + timeout = request.meta.get('download_timeout') or self._connect_timeout + agent = self._get_agent(request, timeout) + + start_time = time() + d = agent.request(request, spider) + d.addCallback(self._cb_latency, request, start_time) + + timeout_cl = reactor.callLater(timeout, d.cancel) + d.addBoth(self._cb_timeout, request, timeout, timeout_cl) + return d + + @staticmethod + def _cb_latency(response: Response, request: Request, start_time: float) -> Response: + request.meta['download_latency'] = time() - start_time + return response + + @staticmethod + def _cb_timeout(response: Response, request: Request, timeout: float, timeout_cl: DelayedCall) -> Response: + if timeout_cl.active(): + timeout_cl.cancel() + return response + + url = urldefrag(request.url)[0] + raise TimeoutError(f"Getting {url} took longer than {timeout} seconds.") diff --git a/scrapy/core/http2/__init__.py b/scrapy/core/http2/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/scrapy/core/http2/agent.py b/scrapy/core/http2/agent.py new file mode 100644 index 000000000..f7b0c3f99 --- /dev/null +++ b/scrapy/core/http2/agent.py @@ -0,0 +1,157 @@ +from collections import deque +from typing import Deque, Dict, List, Optional, Tuple + +from twisted.internet import defer +from twisted.internet.base import ReactorBase +from twisted.internet.defer import Deferred +from twisted.internet.endpoints import HostnameEndpoint +from twisted.python.failure import Failure +from twisted.web.client import URI, BrowserLikePolicyForHTTPS, _StandardEndpointFactory +from twisted.web.error import SchemeNotSupported + +from scrapy.core.downloader.contextfactory import AcceptableProtocolsContextFactory +from scrapy.core.http2.protocol import H2ClientProtocol, H2ClientFactory +from scrapy.http.request import Request +from scrapy.settings import Settings +from scrapy.spiders import Spider + + +class H2ConnectionPool: + def __init__(self, reactor: ReactorBase, settings: Settings) -> None: + self._reactor = reactor + self.settings = settings + + # Store a dictionary which is used to get the respective + # H2ClientProtocolInstance using the key as Tuple(scheme, hostname, port) + self._connections: Dict[Tuple, H2ClientProtocol] = {} + + # Save all requests that arrive before the connection is established + self._pending_requests: Dict[Tuple, Deque[Deferred]] = {} + + def get_connection(self, key: Tuple, uri: URI, endpoint: HostnameEndpoint) -> Deferred: + if key in self._pending_requests: + # Received a request while connecting to remote + # Create a deferred which will fire with the H2ClientProtocol + # instance + d = Deferred() + self._pending_requests[key].append(d) + return d + + # Check if we already have a connection to the remote + conn = self._connections.get(key, None) + if conn: + # Return this connection instance wrapped inside a deferred + return defer.succeed(conn) + + # No connection is established for the given URI + return self._new_connection(key, uri, endpoint) + + def _new_connection(self, key: Tuple, uri: URI, endpoint: HostnameEndpoint) -> Deferred: + self._pending_requests[key] = deque() + + conn_lost_deferred = Deferred() + conn_lost_deferred.addCallback(self._remove_connection, key) + + factory = H2ClientFactory(uri, self.settings, conn_lost_deferred) + conn_d = endpoint.connect(factory) + conn_d.addCallback(self.put_connection, key) + + d = Deferred() + self._pending_requests[key].append(d) + return d + + def put_connection(self, conn: H2ClientProtocol, key: Tuple) -> H2ClientProtocol: + self._connections[key] = conn + + # Now as we have established a proper HTTP/2 connection + # we fire all the deferred's with the connection instance + pending_requests = self._pending_requests.pop(key, None) + while pending_requests: + d = pending_requests.popleft() + d.callback(conn) + + return conn + + def _remove_connection(self, errors: List[BaseException], key: Tuple) -> None: + self._connections.pop(key) + + # Call the errback of all the pending requests for this connection + pending_requests = self._pending_requests.pop(key, None) + while pending_requests: + d = pending_requests.popleft() + d.errback(errors) + + def close_connections(self) -> None: + """Close all the HTTP/2 connections and remove them from pool + + Returns: + Deferred that fires when all connections have been closed + """ + for conn in self._connections.values(): + conn.transport.abortConnection() + + +class H2Agent: + def __init__( + self, + reactor: ReactorBase, + pool: H2ConnectionPool, + context_factory: BrowserLikePolicyForHTTPS = BrowserLikePolicyForHTTPS(), + connect_timeout: Optional[float] = None, + bind_address: Optional[bytes] = None, + ) -> None: + self._reactor = reactor + self._pool = pool + self._context_factory = AcceptableProtocolsContextFactory(context_factory, acceptable_protocols=[b'h2']) + self.endpoint_factory = _StandardEndpointFactory( + self._reactor, self._context_factory, connect_timeout, bind_address + ) + + def get_endpoint(self, uri: URI): + return self.endpoint_factory.endpointForURI(uri) + + def get_key(self, uri: URI) -> Tuple: + """ + Arguments: + uri - URI obtained directly from request URL + """ + return uri.scheme, uri.host, uri.port + + def request(self, request: Request, spider: Spider) -> Deferred: + uri = URI.fromBytes(bytes(request.url, encoding='utf-8')) + try: + endpoint = self.get_endpoint(uri) + except SchemeNotSupported: + return defer.fail(Failure()) + + key = self.get_key(uri) + d = self._pool.get_connection(key, uri, endpoint) + d.addCallback(lambda conn: conn.request(request, spider)) + return d + + +class ScrapyProxyH2Agent(H2Agent): + def __init__( + self, + reactor: ReactorBase, + proxy_uri: URI, + pool: H2ConnectionPool, + context_factory: BrowserLikePolicyForHTTPS = BrowserLikePolicyForHTTPS(), + connect_timeout: Optional[float] = None, + bind_address: Optional[bytes] = None, + ) -> None: + super(ScrapyProxyH2Agent, self).__init__( + reactor=reactor, + pool=pool, + context_factory=context_factory, + connect_timeout=connect_timeout, + bind_address=bind_address, + ) + self._proxy_uri = proxy_uri + + def get_endpoint(self, uri: URI): + return self.endpoint_factory.endpointForURI(self._proxy_uri) + + def get_key(self, uri: URI) -> Tuple: + """We use the proxy uri instead of uri obtained from request url""" + return "http-proxy", self._proxy_uri.host, self._proxy_uri.port diff --git a/scrapy/core/http2/protocol.py b/scrapy/core/http2/protocol.py new file mode 100644 index 000000000..1d150b7ce --- /dev/null +++ b/scrapy/core/http2/protocol.py @@ -0,0 +1,418 @@ +import ipaddress +import itertools +import logging +from collections import deque +from ipaddress import IPv4Address, IPv6Address +from typing import Dict, List, Optional, Union + +from h2.config import H2Configuration +from h2.connection import H2Connection +from h2.errors import ErrorCodes +from h2.events import ( + Event, ConnectionTerminated, DataReceived, ResponseReceived, + SettingsAcknowledged, StreamEnded, StreamReset, UnknownFrameReceived, + WindowUpdated +) +from h2.exceptions import FrameTooLargeError, H2Error +from twisted.internet.defer import Deferred +from twisted.internet.error import TimeoutError +from twisted.internet.interfaces import IHandshakeListener, IProtocolNegotiationFactory +from twisted.internet.protocol import connectionDone, Factory, Protocol +from twisted.internet.ssl import Certificate +from twisted.protocols.policies import TimeoutMixin +from twisted.python.failure import Failure +from twisted.web.client import URI +from zope.interface import implementer + +from scrapy.core.http2.stream import Stream, StreamCloseReason +from scrapy.http import Request +from scrapy.settings import Settings +from scrapy.spiders import Spider + + +logger = logging.getLogger(__name__) + + +PROTOCOL_NAME = b"h2" + + +class InvalidNegotiatedProtocol(H2Error): + + def __init__(self, negotiated_protocol: bytes) -> None: + self.negotiated_protocol = negotiated_protocol + + def __str__(self) -> str: + return (f"Expected {PROTOCOL_NAME!r}, received {self.negotiated_protocol!r}") + + +class RemoteTerminatedConnection(H2Error): + def __init__( + self, + remote_ip_address: Optional[Union[IPv4Address, IPv6Address]], + event: ConnectionTerminated, + ) -> None: + self.remote_ip_address = remote_ip_address + self.terminate_event = event + + def __str__(self) -> str: + return f'Received GOAWAY frame from {self.remote_ip_address!r}' + + +class MethodNotAllowed405(H2Error): + def __init__(self, remote_ip_address: Optional[Union[IPv4Address, IPv6Address]]) -> None: + self.remote_ip_address = remote_ip_address + + def __str__(self) -> str: + return f"Received 'HTTP/2.0 405 Method Not Allowed' from {self.remote_ip_address!r}" + + +@implementer(IHandshakeListener) +class H2ClientProtocol(Protocol, TimeoutMixin): + IDLE_TIMEOUT = 240 + + def __init__(self, uri: URI, settings: Settings, conn_lost_deferred: Deferred) -> None: + """ + Arguments: + uri -- URI of the base url to which HTTP/2 Connection will be made. + uri is used to verify that incoming client requests have correct + base URL. + settings -- Scrapy project settings + conn_lost_deferred -- Deferred fires with the reason: Failure to notify + that connection was lost + """ + self._conn_lost_deferred = conn_lost_deferred + + config = H2Configuration(client_side=True, header_encoding='utf-8') + self.conn = H2Connection(config=config) + + # ID of the next request stream + # Following the convention - 'Streams initiated by a client MUST + # use odd-numbered stream identifiers' (RFC 7540 - Section 5.1.1) + self._stream_id_generator = itertools.count(start=1, step=2) + + # Streams are stored in a dictionary keyed off their stream IDs + self.streams: Dict[int, Stream] = {} + + # If requests are received before connection is made we keep + # all requests in a pool and send them as the connection is made + self._pending_request_stream_pool: deque = deque() + + # Save an instance of errors raised which lead to losing the connection + # We pass these instances to the streams ResponseFailed() failure + self._conn_lost_errors: List[BaseException] = [] + + # Some meta data of this connection + # initialized when connection is successfully made + self.metadata: Dict = { + # Peer certificate instance + 'certificate': None, + + # Address of the server we are connected to which + # is updated when HTTP/2 connection is made successfully + 'ip_address': None, + + # URI of the peer HTTP/2 connection is made + 'uri': uri, + + # Both ip_address and uri are used by the Stream before + # initiating the request to verify that the base address + + # Variables taken from Project Settings + 'default_download_maxsize': settings.getint('DOWNLOAD_MAXSIZE'), + 'default_download_warnsize': settings.getint('DOWNLOAD_WARNSIZE'), + + # Counter to keep track of opened streams. This counter + # is used to make sure that not more than MAX_CONCURRENT_STREAMS + # streams are opened which leads to ProtocolError + # We use simple FIFO policy to handle pending requests + 'active_streams': 0, + + # Flag to keep track if settings were acknowledged by the remote + # This ensures that we have established a HTTP/2 connection + 'settings_acknowledged': False, + } + + @property + def h2_connected(self) -> bool: + """Boolean to keep track of the connection status. + This is used while initiating pending streams to make sure + that we initiate stream only during active HTTP/2 Connection + """ + return bool(self.transport.connected) and self.metadata['settings_acknowledged'] + + @property + def allowed_max_concurrent_streams(self) -> int: + """We keep total two streams for client (sending data) and + server side (receiving data) for a single request. To be safe + we choose the minimum. Since this value can change in event + RemoteSettingsChanged we make variable a property. + """ + return min( + self.conn.local_settings.max_concurrent_streams, + self.conn.remote_settings.max_concurrent_streams + ) + + def _send_pending_requests(self) -> None: + """Initiate all pending requests from the deque following FIFO + We make sure that at any time {allowed_max_concurrent_streams} + streams are active. + """ + while ( + self._pending_request_stream_pool + and self.metadata['active_streams'] < self.allowed_max_concurrent_streams + and self.h2_connected + ): + self.metadata['active_streams'] += 1 + stream = self._pending_request_stream_pool.popleft() + stream.initiate_request() + self._write_to_transport() + + def pop_stream(self, stream_id: int) -> Stream: + """Perform cleanup when a stream is closed + """ + stream = self.streams.pop(stream_id) + self.metadata['active_streams'] -= 1 + self._send_pending_requests() + return stream + + def _new_stream(self, request: Request, spider: Spider) -> Stream: + """Instantiates a new Stream object + """ + stream = Stream( + stream_id=next(self._stream_id_generator), + request=request, + protocol=self, + download_maxsize=getattr(spider, 'download_maxsize', self.metadata['default_download_maxsize']), + download_warnsize=getattr(spider, 'download_warnsize', self.metadata['default_download_warnsize']), + ) + self.streams[stream.stream_id] = stream + return stream + + def _write_to_transport(self) -> None: + """ Write data to the underlying transport connection + from the HTTP2 connection instance if any + """ + # Reset the idle timeout as connection is still actively sending data + self.resetTimeout() + + data = self.conn.data_to_send() + self.transport.write(data) + + def request(self, request: Request, spider: Spider) -> Deferred: + if not isinstance(request, Request): + raise TypeError(f'Expected scrapy.http.Request, received {request.__class__.__qualname__}') + + stream = self._new_stream(request, spider) + d = stream.get_response() + + # Add the stream to the request pool + self._pending_request_stream_pool.append(stream) + + # If we receive a request when connection is idle + # We need to initiate pending requests + self._send_pending_requests() + return d + + def connectionMade(self) -> None: + """Called by Twisted when the connection is established. We can start + sending some data now: we should open with the connection preamble. + """ + # Initialize the timeout + self.setTimeout(self.IDLE_TIMEOUT) + + destination = self.transport.getPeer() + self.metadata['ip_address'] = ipaddress.ip_address(destination.host) + + # Initiate H2 Connection + self.conn.initiate_connection() + self._write_to_transport() + + def _lose_connection_with_error(self, errors: List[BaseException]) -> None: + """Helper function to lose the connection with the error sent as a + reason""" + self._conn_lost_errors += errors + self.transport.loseConnection() + + def handshakeCompleted(self) -> None: + """ + Close the connection if it's not made via the expected protocol + """ + if self.transport.negotiatedProtocol is not None and self.transport.negotiatedProtocol != PROTOCOL_NAME: + # we have not initiated the connection yet, no need to send a GOAWAY frame to the remote peer + self._lose_connection_with_error([InvalidNegotiatedProtocol(self.transport.negotiatedProtocol)]) + + def _check_received_data(self, data: bytes) -> None: + """Checks for edge cases where the connection to remote fails + without raising an appropriate H2Error + + Arguments: + data -- Data received from the remote + """ + if data.startswith(b'HTTP/2.0 405 Method Not Allowed'): + raise MethodNotAllowed405(self.metadata['ip_address']) + + def dataReceived(self, data: bytes) -> None: + # Reset the idle timeout as connection is still actively receiving data + self.resetTimeout() + + try: + self._check_received_data(data) + events = self.conn.receive_data(data) + self._handle_events(events) + except H2Error as e: + if isinstance(e, FrameTooLargeError): + # hyper-h2 does not drop the connection in this scenario, we + # need to abort the connection manually. + self._conn_lost_errors += [e] + self.transport.abortConnection() + return + + # Save this error as ultimately the connection will be dropped + # internally by hyper-h2. Saved error will be passed to all the streams + # closed with the connection. + self._lose_connection_with_error([e]) + finally: + self._write_to_transport() + + def timeoutConnection(self) -> None: + """Called when the connection times out. + We lose the connection with TimeoutError""" + + # Check whether there are open streams. If there are, we're going to + # want to use the error code PROTOCOL_ERROR. If there aren't, use + # NO_ERROR. + if ( + self.conn.open_outbound_streams > 0 + or self.conn.open_inbound_streams > 0 + or self.metadata['active_streams'] > 0 + ): + error_code = ErrorCodes.PROTOCOL_ERROR + else: + error_code = ErrorCodes.NO_ERROR + self.conn.close_connection(error_code=error_code) + self._write_to_transport() + + self._lose_connection_with_error([ + TimeoutError(f"Connection was IDLE for more than {self.IDLE_TIMEOUT}s") + ]) + + def connectionLost(self, reason: Failure = connectionDone) -> None: + """Called by Twisted when the transport connection is lost. + No need to write anything to transport here. + """ + # Cancel the timeout if not done yet + self.setTimeout(None) + + # Notify the connection pool instance such that no new requests are + # sent over current connection + if not reason.check(connectionDone): + self._conn_lost_errors.append(reason) + + self._conn_lost_deferred.callback(self._conn_lost_errors) + + for stream in self.streams.values(): + if stream.metadata['request_sent']: + close_reason = StreamCloseReason.CONNECTION_LOST + else: + close_reason = StreamCloseReason.INACTIVE + stream.close(close_reason, self._conn_lost_errors, from_protocol=True) + + self.metadata['active_streams'] -= len(self.streams) + self.streams.clear() + self._pending_request_stream_pool.clear() + self.conn.close_connection() + + def _handle_events(self, events: List[Event]) -> None: + """Private method which acts as a bridge between the events + received from the HTTP/2 data and IH2EventsHandler + + Arguments: + events -- A list of events that the remote peer triggered by sending data + """ + for event in events: + if isinstance(event, ConnectionTerminated): + self.connection_terminated(event) + elif isinstance(event, DataReceived): + self.data_received(event) + elif isinstance(event, ResponseReceived): + self.response_received(event) + elif isinstance(event, StreamEnded): + self.stream_ended(event) + elif isinstance(event, StreamReset): + self.stream_reset(event) + elif isinstance(event, WindowUpdated): + self.window_updated(event) + elif isinstance(event, SettingsAcknowledged): + self.settings_acknowledged(event) + elif isinstance(event, UnknownFrameReceived): + logger.warning('Unknown frame received: %s', event.frame) + + # Event handler functions starts here + def connection_terminated(self, event: ConnectionTerminated) -> None: + self._lose_connection_with_error([ + RemoteTerminatedConnection(self.metadata['ip_address'], event) + ]) + + def data_received(self, event: DataReceived) -> None: + try: + stream = self.streams[event.stream_id] + except KeyError: + pass # We ignore server-initiated events + else: + stream.receive_data(event.data, event.flow_controlled_length) + + def response_received(self, event: ResponseReceived) -> None: + try: + stream = self.streams[event.stream_id] + except KeyError: + pass # We ignore server-initiated events + else: + stream.receive_headers(event.headers) + + def settings_acknowledged(self, event: SettingsAcknowledged) -> None: + self.metadata['settings_acknowledged'] = True + + # Send off all the pending requests as now we have + # established a proper HTTP/2 connection + self._send_pending_requests() + + # Update certificate when our HTTP/2 connection is established + self.metadata['certificate'] = Certificate(self.transport.getPeerCertificate()) + + def stream_ended(self, event: StreamEnded) -> None: + try: + stream = self.pop_stream(event.stream_id) + except KeyError: + pass # We ignore server-initiated events + else: + stream.close(StreamCloseReason.ENDED, from_protocol=True) + + def stream_reset(self, event: StreamReset) -> None: + try: + stream = self.pop_stream(event.stream_id) + except KeyError: + pass # We ignore server-initiated events + else: + stream.close(StreamCloseReason.RESET, from_protocol=True) + + def window_updated(self, event: WindowUpdated) -> None: + if event.stream_id != 0: + self.streams[event.stream_id].receive_window_update() + else: + # Send leftover data for all the streams + for stream in self.streams.values(): + stream.receive_window_update() + + +@implementer(IProtocolNegotiationFactory) +class H2ClientFactory(Factory): + def __init__(self, uri: URI, settings: Settings, conn_lost_deferred: Deferred) -> None: + self.uri = uri + self.settings = settings + self.conn_lost_deferred = conn_lost_deferred + + def buildProtocol(self, addr) -> H2ClientProtocol: + return H2ClientProtocol(self.uri, self.settings, self.conn_lost_deferred) + + def acceptableProtocols(self) -> List[bytes]: + return [PROTOCOL_NAME] diff --git a/scrapy/core/http2/stream.py b/scrapy/core/http2/stream.py new file mode 100644 index 000000000..c2a4b702f --- /dev/null +++ b/scrapy/core/http2/stream.py @@ -0,0 +1,470 @@ +import logging +from enum import Enum +from io import BytesIO +from urllib.parse import urlparse +from typing import Dict, List, Optional, Tuple, TYPE_CHECKING + +from h2.errors import ErrorCodes +from h2.exceptions import H2Error, ProtocolError, StreamClosedError +from hpack import HeaderTuple +from twisted.internet.defer import Deferred, CancelledError +from twisted.internet.error import ConnectionClosed +from twisted.python.failure import Failure +from twisted.web.client import ResponseFailed + +from scrapy.http import Request +from scrapy.http.headers import Headers +from scrapy.responsetypes import responsetypes + +if TYPE_CHECKING: + from scrapy.core.http2.protocol import H2ClientProtocol + + +logger = logging.getLogger(__name__) + + +class InactiveStreamClosed(ConnectionClosed): + """Connection was closed without sending request headers + of the stream. This happens when a stream is waiting for other + streams to close and connection is lost.""" + + def __init__(self, request: Request) -> None: + self.request = request + + def __str__(self) -> str: + return f'InactiveStreamClosed: Connection was closed without sending the request {self.request!r}' + + +class InvalidHostname(H2Error): + + def __init__(self, request: Request, expected_hostname: str, expected_netloc: str) -> None: + self.request = request + self.expected_hostname = expected_hostname + self.expected_netloc = expected_netloc + + def __str__(self) -> str: + return f'InvalidHostname: Expected {self.expected_hostname} or {self.expected_netloc} in {self.request}' + + +class StreamCloseReason(Enum): + # Received a StreamEnded event from the remote + ENDED = 1 + + # Received a StreamReset event -- ended abruptly + RESET = 2 + + # Transport connection was lost + CONNECTION_LOST = 3 + + # Expected response body size is more than allowed limit + MAXSIZE_EXCEEDED = 4 + + # Response deferred is cancelled by the client + # (happens when client called response_deferred.cancel()) + CANCELLED = 5 + + # Connection lost and the stream was not initiated + INACTIVE = 6 + + # The hostname of the request is not same as of connected peer hostname + # As a result sending this request will the end the connection + INVALID_HOSTNAME = 7 + + +class Stream: + """Represents a single HTTP/2 Stream. + + Stream is a bidirectional flow of bytes within an established connection, + which may carry one or more messages. Handles the transfer of HTTP Headers + and Data frames. + + Role of this class is to + 1. Combine all the data frames + """ + + def __init__( + self, + stream_id: int, + request: Request, + protocol: "H2ClientProtocol", + download_maxsize: int = 0, + download_warnsize: int = 0, + ) -> None: + """ + Arguments: + stream_id -- Unique identifier for the stream within a single HTTP/2 connection + request -- The HTTP request associated to the stream + protocol -- Parent H2ClientProtocol instance + """ + self.stream_id: int = stream_id + self._request: Request = request + self._protocol: "H2ClientProtocol" = protocol + + self._download_maxsize = self._request.meta.get('download_maxsize', download_maxsize) + self._download_warnsize = self._request.meta.get('download_warnsize', download_warnsize) + + # Metadata of an HTTP/2 connection stream + # initialized when stream is instantiated + self.metadata: Dict = { + 'request_content_length': 0 if self._request.body is None else len(self._request.body), + + # Flag to keep track whether the stream has initiated the request + 'request_sent': False, + + # Flag to track whether we have logged about exceeding download warnsize + 'reached_warnsize': False, + + # Each time we send a data frame, we will decrease value by the amount send. + 'remaining_content_length': 0 if self._request.body is None else len(self._request.body), + + # Flag to keep track whether client (self) have closed this stream + 'stream_closed_local': False, + + # Flag to keep track whether the server has closed the stream + 'stream_closed_server': False, + } + + # Private variable used to build the response + # this response is then converted to appropriate Response class + # passed to the response deferred callback + self._response: Dict = { + # Data received frame by frame from the server is appended + # and passed to the response Deferred when completely received. + 'body': BytesIO(), + + # The amount of data received that counts against the + # flow control window + 'flow_controlled_size': 0, + + # Headers received after sending the request + 'headers': Headers({}), + } + + def _cancel(_) -> None: + # Close this stream as gracefully as possible + # If the associated request is initiated we reset this stream + # else we directly call close() method + if self.metadata['request_sent']: + self.reset_stream(StreamCloseReason.CANCELLED) + else: + self.close(StreamCloseReason.CANCELLED) + + self._deferred_response = Deferred(_cancel) + + def __str__(self) -> str: + return f'Stream(id={self.stream_id!r})' + + __repr__ = __str__ + + @property + def _log_warnsize(self) -> bool: + """Checks if we have received data which exceeds the download warnsize + and whether we have not already logged about it. + + Returns: + True if both the above conditions hold true + False if any of the conditions is false + """ + content_length_header = int(self._response['headers'].get(b'Content-Length', -1)) + return ( + self._download_warnsize + and ( + self._response['flow_controlled_size'] > self._download_warnsize + or content_length_header > self._download_warnsize + ) + and not self.metadata['reached_warnsize'] + ) + + def get_response(self) -> Deferred: + """Simply return a Deferred which fires when response + from the asynchronous request is available + """ + return self._deferred_response + + def check_request_url(self) -> bool: + # Make sure that we are sending the request to the correct URL + url = urlparse(self._request.url) + return ( + url.netloc == str(self._protocol.metadata['uri'].host, 'utf-8') + or url.netloc == str(self._protocol.metadata['uri'].netloc, 'utf-8') + or url.netloc == f'{self._protocol.metadata["ip_address"]}:{self._protocol.metadata["uri"].port}' + ) + + def _get_request_headers(self) -> List[Tuple[str, str]]: + url = urlparse(self._request.url) + + path = url.path + if url.query: + path += '?' + url.query + + # This pseudo-header field MUST NOT be empty for "http" or "https" + # URIs; "http" or "https" URIs that do not contain a path component + # MUST include a value of '/'. The exception to this rule is an + # OPTIONS request for an "http" or "https" URI that does not include + # a path component; these MUST include a ":path" pseudo-header field + # with a value of '*' (refer RFC 7540 - Section 8.1.2.3) + if not path: + path = '*' if self._request.method == 'OPTIONS' else '/' + + # Make sure pseudo-headers comes before all the other headers + headers = [ + (':method', self._request.method), + (':authority', url.netloc), + ] + + # The ":scheme" and ":path" pseudo-header fields MUST + # be omitted for CONNECT method (refer RFC 7540 - Section 8.3) + if self._request.method != 'CONNECT': + headers += [ + (':scheme', self._protocol.metadata['uri'].scheme), + (':path', path), + ] + + content_length = str(len(self._request.body)) + headers.append(('Content-Length', content_length)) + + content_length_name = self._request.headers.normkey(b'Content-Length') + for name, values in self._request.headers.items(): + for value in values: + value = str(value, 'utf-8') + if name == content_length_name: + if value != content_length: + logger.warning( + 'Ignoring bad Content-Length header %r of request %r, ' + 'sending %r instead', + value, + self._request, + content_length, + ) + continue + headers.append((str(name, 'utf-8'), value)) + + return headers + + def initiate_request(self) -> None: + if self.check_request_url(): + headers = self._get_request_headers() + self._protocol.conn.send_headers(self.stream_id, headers, end_stream=False) + self.metadata['request_sent'] = True + self.send_data() + else: + # Close this stream calling the response errback + # Note that we have not sent any headers + self.close(StreamCloseReason.INVALID_HOSTNAME) + + def send_data(self) -> None: + """Called immediately after the headers are sent. Here we send all the + data as part of the request. + + If the content length is 0 initially then we end the stream immediately and + wait for response data. + + Warning: Only call this method when stream not closed from client side + and has initiated request already by sending HEADER frame. If not then + stream will raise ProtocolError (raise by h2 state machine). + """ + if self.metadata['stream_closed_local']: + raise StreamClosedError(self.stream_id) + + # Firstly, check what the flow control window is for current stream. + window_size = self._protocol.conn.local_flow_control_window(stream_id=self.stream_id) + + # Next, check what the maximum frame size is. + max_frame_size = self._protocol.conn.max_outbound_frame_size + + # We will send no more than the window size or the remaining file size + # of data in this call, whichever is smaller. + bytes_to_send_size = min(window_size, self.metadata['remaining_content_length']) + + # We now need to send a number of data frames. + while bytes_to_send_size > 0: + chunk_size = min(bytes_to_send_size, max_frame_size) + + data_chunk_start_id = self.metadata['request_content_length'] - self.metadata['remaining_content_length'] + data_chunk = self._request.body[data_chunk_start_id:data_chunk_start_id + chunk_size] + + self._protocol.conn.send_data(self.stream_id, data_chunk, end_stream=False) + + bytes_to_send_size = bytes_to_send_size - chunk_size + self.metadata['remaining_content_length'] = self.metadata['remaining_content_length'] - chunk_size + + self.metadata['remaining_content_length'] = max(0, self.metadata['remaining_content_length']) + + # End the stream if no more data needs to be send + if self.metadata['remaining_content_length'] == 0: + self._protocol.conn.end_stream(self.stream_id) + + # Q. What about the rest of the data? + # Ans: Remaining Data frames will be sent when we get a WindowUpdate frame + + def receive_window_update(self) -> None: + """Flow control window size was changed. + Send data that earlier could not be sent as we were + blocked behind the flow control. + """ + if ( + self.metadata['remaining_content_length'] + and not self.metadata['stream_closed_server'] + and self.metadata['request_sent'] + ): + self.send_data() + + def receive_data(self, data: bytes, flow_controlled_length: int) -> None: + self._response['body'].write(data) + self._response['flow_controlled_size'] += flow_controlled_length + + # We check maxsize here in case the Content-Length header was not received + if self._download_maxsize and self._response['flow_controlled_size'] > self._download_maxsize: + self.reset_stream(StreamCloseReason.MAXSIZE_EXCEEDED) + return + + if self._log_warnsize: + self.metadata['reached_warnsize'] = True + warning_msg = ( + f'Received more ({self._response["flow_controlled_size"]}) bytes than download ' + f'warn size ({self._download_warnsize}) in request {self._request}' + ) + logger.warning(warning_msg) + + # Acknowledge the data received + self._protocol.conn.acknowledge_received_data( + self._response['flow_controlled_size'], + self.stream_id + ) + + def receive_headers(self, headers: List[HeaderTuple]) -> None: + for name, value in headers: + self._response['headers'][name] = value + + # Check if we exceed the allowed max data size which can be received + expected_size = int(self._response['headers'].get(b'Content-Length', -1)) + if self._download_maxsize and expected_size > self._download_maxsize: + self.reset_stream(StreamCloseReason.MAXSIZE_EXCEEDED) + return + + if self._log_warnsize: + self.metadata['reached_warnsize'] = True + warning_msg = ( + f'Expected response size ({expected_size}) larger than ' + f'download warn size ({self._download_warnsize}) in request {self._request}' + ) + logger.warning(warning_msg) + + def reset_stream(self, reason: StreamCloseReason = StreamCloseReason.RESET) -> None: + """Close this stream by sending a RST_FRAME to the remote peer""" + if self.metadata['stream_closed_local']: + raise StreamClosedError(self.stream_id) + + # Clear buffer earlier to avoid keeping data in memory for a long time + self._response['body'].truncate(0) + + self.metadata['stream_closed_local'] = True + self._protocol.conn.reset_stream(self.stream_id, ErrorCodes.REFUSED_STREAM) + self.close(reason) + + def close( + self, + reason: StreamCloseReason, + errors: Optional[List[BaseException]] = None, + from_protocol: bool = False, + ) -> None: + """Based on the reason sent we will handle each case. + """ + if self.metadata['stream_closed_server']: + raise StreamClosedError(self.stream_id) + + if not isinstance(reason, StreamCloseReason): + raise TypeError(f'Expected StreamCloseReason, received {reason.__class__.__qualname__}') + + # Have default value of errors as an empty list as + # some cases can add a list of exceptions + errors = errors or [] + + if not from_protocol: + self._protocol.pop_stream(self.stream_id) + + self.metadata['stream_closed_server'] = True + + # We do not check for Content-Length or Transfer-Encoding in response headers + # and add `partial` flag as in HTTP/1.1 as 'A request or response that includes + # a payload body can include a content-length header field' (RFC 7540 - Section 8.1.2.6) + + # NOTE: Order of handling the events is important here + # As we immediately cancel the request when maxsize is exceeded while + # receiving DATA_FRAME's when we have received the headers (not + # having Content-Length) + if reason is StreamCloseReason.MAXSIZE_EXCEEDED: + expected_size = int(self._response['headers'].get( + b'Content-Length', + self._response['flow_controlled_size']) + ) + error_msg = ( + f'Cancelling download of {self._request.url}: received response ' + f'size ({expected_size}) larger than download max size ({self._download_maxsize})' + ) + logger.error(error_msg) + self._deferred_response.errback(CancelledError(error_msg)) + + elif reason is StreamCloseReason.ENDED: + self._fire_response_deferred() + + # Stream was abruptly ended here + elif reason is StreamCloseReason.CANCELLED: + # Client has cancelled the request. Remove all the data + # received and fire the response deferred with no flags set + + # NOTE: The data is already flushed in Stream.reset_stream() called + # immediately when the stream needs to be cancelled + + # There maybe no :status in headers, we make + # HTTP Status Code: 499 - Client Closed Request + self._response['headers'][':status'] = '499' + self._fire_response_deferred() + + elif reason is StreamCloseReason.RESET: + self._deferred_response.errback(ResponseFailed([ + Failure( + f'Remote peer {self._protocol.metadata["ip_address"]} sent RST_STREAM', + ProtocolError + ) + ])) + + elif reason is StreamCloseReason.CONNECTION_LOST: + self._deferred_response.errback(ResponseFailed(errors)) + + elif reason is StreamCloseReason.INACTIVE: + errors.insert(0, InactiveStreamClosed(self._request)) + self._deferred_response.errback(ResponseFailed(errors)) + + else: + assert reason is StreamCloseReason.INVALID_HOSTNAME + self._deferred_response.errback(InvalidHostname( + self._request, + str(self._protocol.metadata['uri'].host, 'utf-8'), + f'{self._protocol.metadata["ip_address"]}:{self._protocol.metadata["uri"].port}' + )) + + def _fire_response_deferred(self) -> None: + """Builds response from the self._response dict + and fires the response deferred callback with the + generated response instance""" + + body = self._response['body'].getvalue() + response_cls = responsetypes.from_args( + headers=self._response['headers'], + url=self._request.url, + body=body, + ) + + response = response_cls( + url=self._request.url, + status=int(self._response['headers'][':status']), + headers=self._response['headers'], + body=body, + request=self._request, + certificate=self._protocol.metadata['certificate'], + ip_address=self._protocol.metadata['ip_address'], + protocol='h2', + ) + + self._deferred_response.callback(response) diff --git a/scrapy/utils/log.py b/scrapy/utils/log.py index 62df7a6ab..6c456ed60 100644 --- a/scrapy/utils/log.py +++ b/scrapy/utils/log.py @@ -46,6 +46,9 @@ DEFAULT_LOGGING = { 'version': 1, 'disable_existing_loggers': False, 'loggers': { + 'hpack': { + 'level': 'ERROR', + }, 'scrapy': { 'level': 'DEBUG', }, diff --git a/setup.py b/setup.py index cf9261271..767c6f6bf 100644 --- a/setup.py +++ b/setup.py @@ -19,7 +19,7 @@ def has_environment_marker_platform_impl_support(): install_requires = [ - 'Twisted>=17.9.0', + 'Twisted[http2]>=17.9.0', 'cryptography>=2.0', 'cssselect>=0.9.1', 'itemloaders>=1.0.1', @@ -31,6 +31,7 @@ install_requires = [ 'zope.interface>=4.1.3', 'protego>=0.1.15', 'itemadapter>=0.1.0', + 'h2>=3.2.0', ] extras_require = {} cpython_dependencies = [ diff --git a/tests/constraints.txt b/tests/constraints.txt deleted file mode 100644 index 5655ac2d3..000000000 --- a/tests/constraints.txt +++ /dev/null @@ -1 +0,0 @@ -Twisted!=18.4.0 \ No newline at end of file diff --git a/tests/test_dependencies.py b/tests/test_dependencies.py index 93e7311d2..5e63ebffb 100644 --- a/tests/test_dependencies.py +++ b/tests/test_dependencies.py @@ -36,7 +36,7 @@ class ScrapyUtilsTest(unittest.TestCase): ) config_parser = ConfigParser() config_parser.read(tox_config_file_path) - pattern = r'Twisted==([\d.]+)' + pattern = r'Twisted\[http2\]==([\d.]+)' match = re.search(pattern, config_parser['pinned']['deps']) pinned_twisted_version_string = match[1] diff --git a/tests/test_downloader_handlers.py b/tests/test_downloader_handlers.py index f51a6cd3c..86d72772c 100644 --- a/tests/test_downloader_handlers.py +++ b/tests/test_downloader_handlers.py @@ -2,6 +2,7 @@ import contextlib import os import shutil import tempfile +from typing import Optional, Type from unittest import mock from testfixtures import LogCapture @@ -24,7 +25,6 @@ from scrapy.core.downloader.handlers.http import HTTPDownloadHandler from scrapy.core.downloader.handlers.http10 import HTTP10DownloadHandler from scrapy.core.downloader.handlers.http11 import HTTP11DownloadHandler from scrapy.core.downloader.handlers.s3 import S3DownloadHandler - from scrapy.exceptions import NotConfigured, ScrapyDeprecationWarning from scrapy.http import Headers, Request from scrapy.http.response.text import TextResponse @@ -33,7 +33,6 @@ from scrapy.spiders import Spider from scrapy.utils.misc import create_instance from scrapy.utils.python import to_bytes from scrapy.utils.test import get_crawler, skip_if_no_boto - from tests.mockserver import MockServer, ssl_context_factory, Echo from tests.spiders import SingleRequestSpider @@ -132,6 +131,7 @@ class ContentLengthHeaderResource(resource.Resource): A testing resource which renders itself as the value of the Content-Length header from the request. """ + def render(self, request): return request.requestHeaders.getRawHeaders(b"content-length")[0] @@ -143,6 +143,7 @@ class ChunkedResource(resource.Resource): request.write(b"chunked ") request.write(b"content\n") request.finish() + reactor.callLater(0, response) return server.NOT_DONE_YET @@ -156,6 +157,7 @@ class BrokenChunkedResource(resource.Resource): # Disable terminating chunk on finish. request.chunked = False closeConnection(request) + reactor.callLater(0, response) return server.NOT_DONE_YET @@ -187,6 +189,7 @@ class EmptyContentTypeHeaderResource(resource.Resource): A testing resource which renders itself as the value of request body without content-type header in response. """ + def render(self, request): request.setHeader("content-type", "") return request.content.read() @@ -198,14 +201,14 @@ class LargeChunkedFileResource(resource.Resource): for i in range(1024): request.write(b"x" * 1024) request.finish() + reactor.callLater(0, response) return server.NOT_DONE_YET class HttpTestCase(unittest.TestCase): - scheme = 'http' - download_handler_cls = HTTPDownloadHandler + download_handler_cls: Type = HTTPDownloadHandler # only used for HTTPS tests keyfile = 'keys/localhost.key' @@ -233,8 +236,10 @@ class HttpTestCase(unittest.TestCase): self.wrapper = WrappingFactory(self.site) self.host = 'localhost' if self.scheme == 'https': + # Using WrappingFactory do not enable HTTP/2 failing all the + # tests with H2DownloadHandler self.port = reactor.listenSSL( - 0, self.wrapper, ssl_context_factory(self.keyfile, self.certfile), + 0, self.site, ssl_context_factory(self.keyfile, self.certfile), interface=self.host) else: self.port = reactor.listenTCP(0, self.wrapper, interface=self.host) @@ -284,7 +289,7 @@ class HttpTestCase(unittest.TestCase): def test_timeout_download_from_spider_nodata_rcvd(self): # client connects but no data is received spider = Spider('foo') - meta = {'download_timeout': 0.2} + meta = {'download_timeout': 0.5} request = Request(self.getURL('wait'), meta=meta) d = self.download_request(request, spider) yield self.assertFailure(d, defer.TimeoutError, error.TimeoutError) @@ -293,7 +298,7 @@ class HttpTestCase(unittest.TestCase): def test_timeout_download_from_spider_server_hangs(self): # client connects, server send headers and some body bytes but hangs spider = Spider('foo') - meta = {'download_timeout': 0.2} + meta = {'download_timeout': 0.5} request = Request(self.getURL('hang-after-headers'), meta=meta) d = self.download_request(request, spider) yield self.assertFailure(d, defer.TimeoutError, error.TimeoutError) @@ -308,16 +313,18 @@ class HttpTestCase(unittest.TestCase): return self.download_request(request, Spider('foo')).addCallback(_test) def test_host_header_seted_in_request_headers(self): - def _test(response): - self.assertEqual(response.body, b'example.com') - self.assertEqual(request.headers.get('Host'), b'example.com') + host = self.host + ':' + str(self.portno) - request = Request(self.getURL('host'), headers={'Host': 'example.com'}) + def _test(response): + self.assertEqual(response.body, host.encode()) + self.assertEqual(request.headers.get('Host'), host.encode()) + + request = Request(self.getURL('host'), headers={'Host': host}) return self.download_request(request, Spider('foo')).addCallback(_test) d = self.download_request(request, Spider('foo')) d.addCallback(lambda r: r.body) - d.addCallback(self.assertEqual, b'example.com') + d.addCallback(self.assertEqual, b'localhost') return d def test_content_length_zero_bodyless_post_request_headers(self): @@ -331,10 +338,11 @@ class HttpTestCase(unittest.TestCase): https://github.com/kennethreitz/requests/issues/405 https://bugs.python.org/issue14721 """ + def _test(response): self.assertEqual(response.body, b'0') - request = Request(self.getURL('contentlength'), method='POST', headers={'Host': 'example.com'}) + request = Request(self.getURL('contentlength'), method='POST') return self.download_request(request, Spider('foo')).addCallback(_test) def test_content_length_zero_bodyless_post_only_one(self): @@ -359,7 +367,7 @@ class HttpTestCase(unittest.TestCase): class Http10TestCase(HttpTestCase): """HTTP 1.0 test case""" - download_handler_cls = HTTP10DownloadHandler + download_handler_cls: Type = HTTP10DownloadHandler def test_protocol(self): request = Request(self.getURL("host"), method="GET") @@ -375,7 +383,7 @@ class Https10TestCase(Http10TestCase): class Http11TestCase(HttpTestCase): """HTTP 1.1 test case""" - download_handler_cls = HTTP11DownloadHandler + download_handler_cls: Type = HTTP11DownloadHandler def test_download_without_maxsize_limit(self): request = Request(self.getURL('file')) @@ -569,7 +577,7 @@ class Https11InvalidDNSPattern(Https11TestCase): class Https11CustomCiphers(unittest.TestCase): scheme = 'https' - download_handler_cls = HTTP11DownloadHandler + download_handler_cls: Type = HTTP11DownloadHandler keyfile = 'keys/localhost.key' certfile = 'keys/localhost.crt' @@ -580,10 +588,9 @@ class Https11CustomCiphers(unittest.TestCase): FilePath(self.tmpname).child("file").setContent(b"0123456789") r = static.File(self.tmpname) self.site = server.Site(r, timeout=None) - self.wrapper = WrappingFactory(self.site) self.host = 'localhost' self.port = reactor.listenSSL( - 0, self.wrapper, ssl_context_factory(self.keyfile, self.certfile, cipher_string='CAMELLIA256-SHA'), + 0, self.site, ssl_context_factory(self.keyfile, self.certfile, cipher_string='CAMELLIA256-SHA'), interface=self.host) self.portno = self.port.getHost().port crawler = get_crawler(settings_dict={'DOWNLOADER_CLIENT_TLS_CIPHERS': 'CAMELLIA256-SHA'}) @@ -610,6 +617,7 @@ class Https11CustomCiphers(unittest.TestCase): class Http11MockServerTestCase(unittest.TestCase): """HTTP 1.1 test case with MockServer""" + settings_dict: Optional[dict] = None def setUp(self): self.mockserver = MockServer() @@ -620,7 +628,7 @@ class Http11MockServerTestCase(unittest.TestCase): @defer.inlineCallbacks def test_download_with_content_length(self): - crawler = get_crawler(SingleRequestSpider) + crawler = get_crawler(SingleRequestSpider, self.settings_dict) # http://localhost:8998/partial set Content-Length to 1024, use download_maxsize= 1000 to avoid # download it yield crawler.crawl(seed=Request(url=self.mockserver.url('/partial'), meta={'download_maxsize': 1000})) @@ -629,7 +637,7 @@ class Http11MockServerTestCase(unittest.TestCase): @defer.inlineCallbacks def test_download(self): - crawler = get_crawler(SingleRequestSpider) + crawler = get_crawler(SingleRequestSpider, self.settings_dict) yield crawler.crawl(seed=Request(url=self.mockserver.url(''))) failure = crawler.spider.meta.get('failure') self.assertTrue(failure is None) @@ -638,7 +646,7 @@ class Http11MockServerTestCase(unittest.TestCase): @defer.inlineCallbacks def test_download_gzip_response(self): - crawler = get_crawler(SingleRequestSpider) + crawler = get_crawler(SingleRequestSpider, self.settings_dict) body = b'1' * 100 # PayloadResource requires body length to be 100 request = Request(self.mockserver.url('/payload'), method='POST', body=body, meta={'download_maxsize': 50}) @@ -676,7 +684,8 @@ class UriResource(resource.Resource): class HttpProxyTestCase(unittest.TestCase): - download_handler_cls = HTTPDownloadHandler + download_handler_cls: Type = HTTPDownloadHandler + expected_http_proxy_request_body = b'http://example.com' def setUp(self): site = server.Site(UriResource(), timeout=None) @@ -699,7 +708,7 @@ class HttpProxyTestCase(unittest.TestCase): def _test(response): self.assertEqual(response.status, 200) self.assertEqual(response.url, request.url) - self.assertEqual(response.body, b'http://example.com') + self.assertEqual(response.body, self.expected_http_proxy_request_body) http_proxy = self.getURL('') request = Request('http://example.com', meta={'proxy': http_proxy}) @@ -728,14 +737,14 @@ class HttpProxyTestCase(unittest.TestCase): class Http10ProxyTestCase(HttpProxyTestCase): - download_handler_cls = HTTP10DownloadHandler + download_handler_cls: Type = HTTP10DownloadHandler def test_download_with_proxy_https_noconnect(self): raise unittest.SkipTest('noconnect is not supported in HTTP10DownloadHandler') class Http11ProxyTestCase(HttpProxyTestCase): - download_handler_cls = HTTP11DownloadHandler + download_handler_cls: Type = HTTP11DownloadHandler @defer.inlineCallbacks def test_download_with_proxy_https_timeout(self): @@ -783,7 +792,7 @@ class S3AnonTestCase(unittest.TestCase): class S3TestCase(unittest.TestCase): - download_handler_cls = S3DownloadHandler + download_handler_cls: Type = S3DownloadHandler # test use same example keys than amazon developer guide # http://s3.amazonaws.com/awsdocs/S3/20060301/s3-dg-20060301.pdf @@ -923,7 +932,6 @@ class S3TestCase(unittest.TestCase): class BaseFTPTestCase(unittest.TestCase): - username = "scrapy" password = "passwd" req_meta = {"ftp_user": username, "ftp_password": password} @@ -961,6 +969,7 @@ class BaseFTPTestCase(unittest.TestCase): def _clean(data): self.download_handler.client.transport.loseConnection() return data + deferred.addCallback(_clean) if callback: deferred.addCallback(callback) @@ -991,6 +1000,7 @@ class BaseFTPTestCase(unittest.TestCase): self.assertEqual(r.status, 200) self.assertEqual(r.body, b'Moooooooooo power!') self.assertEqual(r.headers, {b'Local Filename': [b''], b'Size': [b'18']}) + return self._add_test_callbacks(d, _test) def test_ftp_download_notexist(self): @@ -1000,6 +1010,7 @@ class BaseFTPTestCase(unittest.TestCase): def _test(r): self.assertEqual(r.status, 404) + return self._add_test_callbacks(d, _test) def test_ftp_local_filename(self): @@ -1020,6 +1031,7 @@ class BaseFTPTestCase(unittest.TestCase): with open(local_fname, "rb") as f: self.assertEqual(f.read(), b"I have the power!") os.remove(local_fname) + return self._add_test_callbacks(d, _test) @@ -1036,11 +1048,11 @@ class FTPTestCase(BaseFTPTestCase): def _test(r): self.assertEqual(r.type, ConnectionLost) + return self._add_test_callbacks(d, errback=_test) class AnonymousFTPTestCase(BaseFTPTestCase): - username = "anonymous" req_meta = {} diff --git a/tests/test_downloader_handlers_http2.py b/tests/test_downloader_handlers_http2.py new file mode 100644 index 000000000..439778014 --- /dev/null +++ b/tests/test_downloader_handlers_http2.py @@ -0,0 +1,247 @@ +import json +from unittest import mock + +from pytest import mark +from testfixtures import LogCapture +from twisted.internet import defer, error, reactor +from twisted.trial import unittest +from twisted.web import server +from twisted.web.error import SchemeNotSupported + +from scrapy.core.downloader.handlers.http2 import H2DownloadHandler +from scrapy.http import Request +from scrapy.spiders import Spider +from scrapy.utils.misc import create_instance +from scrapy.utils.test import get_crawler +from tests.mockserver import ssl_context_factory +from tests.test_downloader_handlers import ( + Https11TestCase, Https11CustomCiphers, + Http11MockServerTestCase, Http11ProxyTestCase, + UriResource +) + + +class Https2TestCase(Https11TestCase): + scheme = 'https' + download_handler_cls = H2DownloadHandler + HTTP2_DATALOSS_SKIP_REASON = "Content-Length mismatch raises InvalidBodyLengthError" + + def test_protocol(self): + request = Request(self.getURL("host"), method="GET") + d = self.download_request(request, Spider("foo")) + d.addCallback(lambda r: r.protocol) + d.addCallback(self.assertEqual, "h2") + return d + + @defer.inlineCallbacks + def test_download_with_maxsize_very_large_file(self): + with mock.patch('scrapy.core.http2.stream.logger') as logger: + request = Request(self.getURL('largechunkedfile')) + + def check(logger): + logger.error.assert_called_once_with(mock.ANY) + + d = self.download_request(request, Spider('foo', download_maxsize=1500)) + yield self.assertFailure(d, defer.CancelledError, error.ConnectionAborted) + + # As the error message is logged in the dataReceived callback, we + # have to give a bit of time to the reactor to process the queue + # after closing the connection. + d = defer.Deferred() + d.addCallback(check) + reactor.callLater(.1, d.callback, logger) + yield d + + @defer.inlineCallbacks + def test_unsupported_scheme(self): + request = Request("ftp://unsupported.scheme") + d = self.download_request(request, Spider("foo")) + yield self.assertFailure(d, SchemeNotSupported) + + def test_download_broken_content_cause_data_loss(self, url='broken'): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_download_broken_chunked_content_cause_data_loss(self): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_download_broken_content_allow_data_loss(self, url='broken'): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_download_broken_chunked_content_allow_data_loss(self): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_download_broken_content_allow_data_loss_via_setting(self, url='broken'): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_download_broken_chunked_content_allow_data_loss_via_setting(self): + raise unittest.SkipTest(self.HTTP2_DATALOSS_SKIP_REASON) + + def test_concurrent_requests_same_domain(self): + spider = Spider('foo') + + request1 = Request(self.getURL('file')) + d1 = self.download_request(request1, spider) + d1.addCallback(lambda r: r.body) + d1.addCallback(self.assertEqual, b"0123456789") + + request2 = Request(self.getURL('echo'), method='POST') + d2 = self.download_request(request2, spider) + d2.addCallback(lambda r: r.headers['Content-Length']) + d2.addCallback(self.assertEqual, b"79") + + return defer.DeferredList([d1, d2]) + + @mark.xfail(reason="https://github.com/python-hyper/h2/issues/1247") + def test_connect_request(self): + request = Request(self.getURL('file'), method='CONNECT') + d = self.download_request(request, Spider('foo')) + d.addCallback(lambda r: r.body) + d.addCallback(self.assertEqual, b'') + return d + + def test_custom_content_length_good(self): + request = Request(self.getURL('contentlength')) + custom_content_length = str(len(request.body)) + request.headers['Content-Length'] = custom_content_length + d = self.download_request(request, Spider('foo')) + d.addCallback(lambda r: r.text) + d.addCallback(self.assertEqual, custom_content_length) + return d + + def test_custom_content_length_bad(self): + request = Request(self.getURL('contentlength')) + actual_content_length = str(len(request.body)) + bad_content_length = str(len(request.body) + 1) + request.headers['Content-Length'] = bad_content_length + log = LogCapture() + d = self.download_request(request, Spider('foo')) + d.addCallback(lambda r: r.text) + d.addCallback(self.assertEqual, actual_content_length) + d.addCallback( + lambda _: log.check_present( + ( + 'scrapy.core.http2.stream', + 'WARNING', + f'Ignoring bad Content-Length header ' + f'{bad_content_length!r} of request {request}, sending ' + f'{actual_content_length!r} instead', + ) + ) + ) + d.addCallback( + lambda _: log.uninstall() + ) + return d + + def test_duplicate_header(self): + request = Request(self.getURL('echo')) + header, value1, value2 = 'Custom-Header', 'foo', 'bar' + request.headers.appendlist(header, value1) + request.headers.appendlist(header, value2) + d = self.download_request(request, Spider('foo')) + d.addCallback(lambda r: json.loads(r.text)['headers'][header]) + d.addCallback(self.assertEqual, [value1, value2]) + return d + + +class Https2WrongHostnameTestCase(Https2TestCase): + tls_log_message = ( + 'SSL connection certificate: issuer "/C=XW/ST=XW/L=The ' + 'Internet/O=Scrapy/CN=www.example.com/emailAddress=test@example.com", ' + 'subject "/C=XW/ST=XW/L=The ' + 'Internet/O=Scrapy/CN=www.example.com/emailAddress=test@example.com"' + ) + + # above tests use a server certificate for "localhost", + # client connection to "localhost" too. + # here we test that even if the server certificate is for another domain, + # "www.example.com" in this case, + # the tests still pass + keyfile = 'keys/example-com.key.pem' + certfile = 'keys/example-com.cert.pem' + + +class Https2InvalidDNSId(Https2TestCase): + """Connect to HTTPS hosts with IP while certificate uses domain names IDs.""" + + def setUp(self): + super(Https2InvalidDNSId, self).setUp() + self.host = '127.0.0.1' + + +class Https2InvalidDNSPattern(Https2TestCase): + """Connect to HTTPS hosts where the certificate are issued to an ip instead of a domain.""" + + keyfile = 'keys/localhost.ip.key' + certfile = 'keys/localhost.ip.crt' + + def setUp(self): + try: + from service_identity.exceptions import CertificateError # noqa: F401 + except ImportError: + raise unittest.SkipTest("cryptography lib is too old") + self.tls_log_message = ( + 'SSL connection certificate: issuer "/C=IE/O=Scrapy/CN=127.0.0.1", ' + 'subject "/C=IE/O=Scrapy/CN=127.0.0.1"' + ) + super(Https2InvalidDNSPattern, self).setUp() + + +class Https2CustomCiphers(Https11CustomCiphers): + scheme = 'https' + download_handler_cls = H2DownloadHandler + + +class Http2MockServerTestCase(Http11MockServerTestCase): + """HTTP 2.0 test case with MockServer""" + settings_dict = { + 'DOWNLOAD_HANDLERS': { + 'https': 'scrapy.core.downloader.handlers.http2.H2DownloadHandler' + } + } + + +class Https2ProxyTestCase(Http11ProxyTestCase): + # only used for HTTPS tests + keyfile = 'keys/localhost.key' + certfile = 'keys/localhost.crt' + + scheme = 'https' + host = u'127.0.0.1' + + download_handler_cls = H2DownloadHandler + expected_http_proxy_request_body = b'/' + + def setUp(self): + site = server.Site(UriResource(), timeout=None) + self.port = reactor.listenSSL( + 0, site, + ssl_context_factory(self.keyfile, self.certfile), + interface=self.host + ) + self.portno = self.port.getHost().port + self.download_handler = create_instance(self.download_handler_cls, None, get_crawler()) + self.download_request = self.download_handler.download_request + + def getURL(self, path): + return f"{self.scheme}://{self.host}:{self.portno}/{path}" + + def test_download_with_proxy_https_noconnect(self): + def _test(response): + self.assertEqual(response.status, 200) + self.assertEqual(response.url, request.url) + self.assertEqual(response.body, b'/') + + http_proxy = '%s?noconnect' % self.getURL('') + request = Request('https://example.com', meta={'proxy': http_proxy}) + with self.assertWarnsRegex( + Warning, + r'Using HTTPS proxies in the noconnect mode is not supported by the ' + r'downloader handler.' + ): + return self.download_request(request, Spider('foo')).addCallback(_test) + + @defer.inlineCallbacks + def test_download_with_proxy_https_timeout(self): + with self.assertRaises(NotImplementedError): + yield super(Https2ProxyTestCase, self).test_download_with_proxy_https_timeout() diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py new file mode 100644 index 000000000..8b2f6a11d --- /dev/null +++ b/tests/test_http2_client_protocol.py @@ -0,0 +1,666 @@ +import json +import os +import random +import re +import shutil +import string +from ipaddress import IPv4Address +from unittest import mock +from urllib.parse import urlencode + +from h2.exceptions import InvalidBodyLengthError +from twisted.internet import reactor +from twisted.internet.defer import CancelledError, Deferred, DeferredList, inlineCallbacks +from twisted.internet.endpoints import SSL4ClientEndpoint, SSL4ServerEndpoint +from twisted.internet.error import TimeoutError +from twisted.internet.ssl import optionsForClientTLS, PrivateCertificate, Certificate +from twisted.python.failure import Failure +from twisted.trial.unittest import TestCase +from twisted.web.client import ResponseFailed, URI +from twisted.web.http import Request as TxRequest +from twisted.web.server import Site, NOT_DONE_YET +from twisted.web.static import File + +from scrapy.core.http2.protocol import H2ClientFactory, H2ClientProtocol +from scrapy.core.http2.stream import InactiveStreamClosed, InvalidHostname +from scrapy.http import Request, Response, JsonRequest +from scrapy.settings import Settings +from scrapy.spiders import Spider +from tests.mockserver import ssl_context_factory, LeafResource, Status + + +def generate_random_string(size): + return ''.join(random.choices( + string.ascii_uppercase + string.digits, + k=size + )) + + +def make_html_body(val): + response = f''' +

Hello from HTTP2

+

{val}

+''' + return bytes(response, 'utf-8') + + +class DummySpider(Spider): + name = 'dummy' + start_urls: list = [] + + def parse(self, response): + print(response) + + +class Data: + SMALL_SIZE = 1024 # 1 KB + LARGE_SIZE = 1024 ** 2 # 1 MB + + STR_SMALL = generate_random_string(SMALL_SIZE) + STR_LARGE = generate_random_string(LARGE_SIZE) + + EXTRA_SMALL = generate_random_string(1024 * 15) + EXTRA_LARGE = generate_random_string((1024 ** 2) * 15) + + HTML_SMALL = make_html_body(STR_SMALL) + HTML_LARGE = make_html_body(STR_LARGE) + + JSON_SMALL = {'data': STR_SMALL} + JSON_LARGE = {'data': STR_LARGE} + + DATALOSS = b'Dataloss Content' + NO_CONTENT_LENGTH = b'This response do not have any content-length header' + + +class GetDataHtmlSmall(LeafResource): + def render_GET(self, request: TxRequest): + request.setHeader('Content-Type', 'text/html; charset=UTF-8') + return Data.HTML_SMALL + + +class GetDataHtmlLarge(LeafResource): + def render_GET(self, request: TxRequest): + request.setHeader('Content-Type', 'text/html; charset=UTF-8') + return Data.HTML_LARGE + + +class PostDataJsonMixin: + @staticmethod + def make_response(request: TxRequest, extra_data: str): + response = { + 'request-headers': {}, + 'request-body': json.loads(request.content.read()), + 'extra-data': extra_data + } + for k, v in request.requestHeaders.getAllRawHeaders(): + response['request-headers'][str(k, 'utf-8')] = str(v[0], 'utf-8') + + response_bytes = bytes(json.dumps(response), 'utf-8') + request.setHeader('Content-Type', 'application/json; charset=UTF-8') + request.setHeader('Content-Encoding', 'UTF-8') + return response_bytes + + +class PostDataJsonSmall(LeafResource, PostDataJsonMixin): + def render_POST(self, request: TxRequest): + return self.make_response(request, Data.EXTRA_SMALL) + + +class PostDataJsonLarge(LeafResource, PostDataJsonMixin): + def render_POST(self, request: TxRequest): + return self.make_response(request, Data.EXTRA_LARGE) + + +class Dataloss(LeafResource): + + def render_GET(self, request: TxRequest): + request.setHeader(b"Content-Length", b"1024") + self.deferRequest(request, 0, self._delayed_render, request) + return NOT_DONE_YET + + @staticmethod + def _delayed_render(request: TxRequest): + request.write(Data.DATALOSS) + request.finish() + + +class NoContentLengthHeader(LeafResource): + def render_GET(self, request: TxRequest): + request.requestHeaders.removeHeader('Content-Length') + self.deferRequest(request, 0, self._delayed_render, request) + return NOT_DONE_YET + + @staticmethod + def _delayed_render(request: TxRequest): + request.write(Data.NO_CONTENT_LENGTH) + request.finish() + + +class TimeoutResponse(LeafResource): + def render_GET(self, request: TxRequest): + return NOT_DONE_YET + + +class QueryParams(LeafResource): + def render_GET(self, request: TxRequest): + request.setHeader('Content-Type', 'application/json; charset=UTF-8') + request.setHeader('Content-Encoding', 'UTF-8') + + query_params = {} + for k, v in request.args.items(): + query_params[str(k, 'utf-8')] = str(v[0], 'utf-8') + + return bytes(json.dumps(query_params), 'utf-8') + + +class RequestHeaders(LeafResource): + """Sends all the headers received as a response""" + + def render_GET(self, request: TxRequest): + request.setHeader('Content-Type', 'application/json; charset=UTF-8') + request.setHeader('Content-Encoding', 'UTF-8') + headers = {} + for k, v in request.requestHeaders.getAllRawHeaders(): + headers[str(k, 'utf-8')] = str(v[0], 'utf-8') + + return bytes(json.dumps(headers), 'utf-8') + + +def get_client_certificate(key_file, certificate_file) -> PrivateCertificate: + with open(key_file, 'r') as key, open(certificate_file, 'r') as certificate: + pem = ''.join(key.readlines()) + ''.join(certificate.readlines()) + + return PrivateCertificate.loadPEM(pem) + + +class Https2ClientProtocolTestCase(TestCase): + scheme = 'https' + key_file = os.path.join(os.path.dirname(__file__), 'keys', 'localhost.key') + certificate_file = os.path.join(os.path.dirname(__file__), 'keys', 'localhost.crt') + + def _init_resource(self): + self.temp_directory = self.mktemp() + os.mkdir(self.temp_directory) + r = File(self.temp_directory) + r.putChild(b'get-data-html-small', GetDataHtmlSmall()) + r.putChild(b'get-data-html-large', GetDataHtmlLarge()) + + r.putChild(b'post-data-json-small', PostDataJsonSmall()) + r.putChild(b'post-data-json-large', PostDataJsonLarge()) + + r.putChild(b'dataloss', Dataloss()) + r.putChild(b'no-content-length-header', NoContentLengthHeader()) + r.putChild(b'status', Status()) + r.putChild(b'query-params', QueryParams()) + r.putChild(b'timeout', TimeoutResponse()) + r.putChild(b'request-headers', RequestHeaders()) + return r + + @inlineCallbacks + def setUp(self): + # Initialize resource tree + root = self._init_resource() + self.site = Site(root, timeout=None) + + # Start server for testing + self.hostname = u'localhost' + context_factory = ssl_context_factory(self.key_file, self.certificate_file) + + server_endpoint = SSL4ServerEndpoint(reactor, 0, context_factory, interface=self.hostname) + self.server = yield server_endpoint.listen(self.site) + self.port_number = self.server.getHost().port + + # Connect H2 client with server + self.client_certificate = get_client_certificate(self.key_file, self.certificate_file) + client_options = optionsForClientTLS( + hostname=self.hostname, + trustRoot=self.client_certificate, + acceptableProtocols=[b'h2'] + ) + uri = URI.fromBytes(bytes(self.get_url('/'), 'utf-8')) + + self.conn_closed_deferred = Deferred() + h2_client_factory = H2ClientFactory(uri, Settings(), self.conn_closed_deferred) + client_endpoint = SSL4ClientEndpoint(reactor, self.hostname, self.port_number, client_options) + self.client = yield client_endpoint.connect(h2_client_factory) + + @inlineCallbacks + def tearDown(self): + if self.client.connected: + yield self.client.transport.loseConnection() + yield self.client.transport.abortConnection() + yield self.server.stopListening() + shutil.rmtree(self.temp_directory) + self.conn_closed_deferred = None + + def get_url(self, path): + """ + :param path: Should have / at the starting compulsorily if not empty + :return: Complete url + """ + assert len(path) > 0 and (path[0] == '/' or path[0] == '&') + return f'{self.scheme}://{self.hostname}:{self.port_number}{path}' + + def make_request(self, request: Request) -> Deferred: + return self.client.request(request, DummySpider()) + + @staticmethod + def _check_repeat(get_deferred, count): + d_list = [] + for _ in range(count): + d = get_deferred() + d_list.append(d) + + return DeferredList(d_list, fireOnOneErrback=True) + + def _check_GET( + self, + request: Request, + expected_body, + expected_status + ): + def check_response(response: Response): + self.assertEqual(response.status, expected_status) + self.assertEqual(response.body, expected_body) + self.assertEqual(response.request, request) + + content_length = int(response.headers.get('Content-Length')) + self.assertEqual(len(response.body), content_length) + + d = self.make_request(request) + d.addCallback(check_response) + d.addErrback(self.fail) + return d + + def test_GET_small_body(self): + request = Request(self.get_url('/get-data-html-small')) + return self._check_GET(request, Data.HTML_SMALL, 200) + + def test_GET_large_body(self): + request = Request(self.get_url('/get-data-html-large')) + return self._check_GET(request, Data.HTML_LARGE, 200) + + def _check_GET_x10(self, *args, **kwargs): + def get_deferred(): + return self._check_GET(*args, **kwargs) + + return self._check_repeat(get_deferred, 10) + + def test_GET_small_body_x10(self): + return self._check_GET_x10( + Request(self.get_url('/get-data-html-small')), + Data.HTML_SMALL, + 200 + ) + + def test_GET_large_body_x10(self): + return self._check_GET_x10( + Request(self.get_url('/get-data-html-large')), + Data.HTML_LARGE, + 200 + ) + + def _check_POST_json( + self, + request: Request, + expected_request_body, + expected_extra_data, + expected_status: int + ): + d = self.make_request(request) + + def assert_response(response: Response): + self.assertEqual(response.status, expected_status) + self.assertEqual(response.request, request) + + content_length = int(response.headers.get('Content-Length')) + self.assertEqual(len(response.body), content_length) + + # Parse the body + content_encoding = str(response.headers[b'Content-Encoding'], 'utf-8') + body = json.loads(str(response.body, content_encoding)) + self.assertIn('request-body', body) + self.assertIn('extra-data', body) + self.assertIn('request-headers', body) + + request_body = body['request-body'] + self.assertEqual(request_body, expected_request_body) + + extra_data = body['extra-data'] + self.assertEqual(extra_data, expected_extra_data) + + # Check if headers were sent successfully + request_headers = body['request-headers'] + for k, v in request.headers.items(): + k_str = str(k, 'utf-8') + self.assertIn(k_str, request_headers) + self.assertEqual(request_headers[k_str], str(v[0], 'utf-8')) + + d.addCallback(assert_response) + d.addErrback(self.fail) + return d + + def test_POST_small_json(self): + request = JsonRequest(url=self.get_url('/post-data-json-small'), method='POST', data=Data.JSON_SMALL) + return self._check_POST_json( + request, + Data.JSON_SMALL, + Data.EXTRA_SMALL, + 200 + ) + + def test_POST_large_json(self): + request = JsonRequest(url=self.get_url('/post-data-json-large'), method='POST', data=Data.JSON_LARGE) + return self._check_POST_json( + request, + Data.JSON_LARGE, + Data.EXTRA_LARGE, + 200 + ) + + def _check_POST_json_x10(self, *args, **kwargs): + def get_deferred(): + return self._check_POST_json(*args, **kwargs) + + return self._check_repeat(get_deferred, 10) + + def test_POST_small_json_x10(self): + request = JsonRequest(url=self.get_url('/post-data-json-small'), method='POST', data=Data.JSON_SMALL) + return self._check_POST_json_x10( + request, + Data.JSON_SMALL, + Data.EXTRA_SMALL, + 200 + ) + + def test_POST_large_json_x10(self): + request = JsonRequest(url=self.get_url('/post-data-json-large'), method='POST', data=Data.JSON_LARGE) + return self._check_POST_json_x10( + request, + Data.JSON_LARGE, + Data.EXTRA_LARGE, + 200 + ) + + @inlineCallbacks + def test_invalid_negotiated_protocol(self): + with mock.patch("scrapy.core.http2.protocol.PROTOCOL_NAME", return_value=b"not-h2"): + request = Request(url=self.get_url('/status?n=200')) + with self.assertRaises(ResponseFailed): + yield self.make_request(request) + + def test_cancel_request(self): + request = Request(url=self.get_url('/get-data-html-large')) + + def assert_response(response: Response): + self.assertEqual(response.status, 499) + self.assertEqual(response.request, request) + + d = self.make_request(request) + d.addCallback(assert_response) + d.addErrback(self.fail) + d.cancel() + + return d + + def test_download_maxsize_exceeded(self): + request = Request(url=self.get_url('/get-data-html-large'), meta={'download_maxsize': 1000}) + + def assert_cancelled_error(failure): + self.assertIsInstance(failure.value, CancelledError) + error_pattern = re.compile( + rf'Cancelling download of {request.url}: received response ' + rf'size \(\d*\) larger than download max size \(1000\)' + ) + self.assertEqual(len(re.findall(error_pattern, str(failure.value))), 1) + + d = self.make_request(request) + d.addCallback(self.fail) + d.addErrback(assert_cancelled_error) + return d + + def test_received_dataloss_response(self): + """In case when value of Header Content-Length != len(Received Data) + ProtocolError is raised""" + request = Request(url=self.get_url('/dataloss')) + + def assert_failure(failure: Failure): + self.assertTrue(len(failure.value.reasons) > 0) + self.assertTrue(any( + isinstance(error, InvalidBodyLengthError) + for error in failure.value.reasons + )) + + d = self.make_request(request) + d.addCallback(self.fail) + d.addErrback(assert_failure) + return d + + def test_missing_content_length_header(self): + request = Request(url=self.get_url('/no-content-length-header')) + + def assert_content_length(response: Response): + self.assertEqual(response.status, 200) + self.assertEqual(response.body, Data.NO_CONTENT_LENGTH) + self.assertEqual(response.request, request) + self.assertNotIn('Content-Length', response.headers) + + d = self.make_request(request) + d.addCallback(assert_content_length) + d.addErrback(self.fail) + return d + + @inlineCallbacks + def _check_log_warnsize( + self, + request, + warn_pattern, + expected_body + ): + with self.assertLogs('scrapy.core.http2.stream', level='WARNING') as cm: + response = yield self.make_request(request) + self.assertEqual(response.status, 200) + self.assertEqual(response.request, request) + self.assertEqual(response.body, expected_body) + + # Check the warning is raised only once for this request + self.assertEqual(sum( + len(re.findall(warn_pattern, log)) + for log in cm.output + ), 1) + + @inlineCallbacks + def test_log_expected_warnsize(self): + request = Request(url=self.get_url('/get-data-html-large'), meta={'download_warnsize': 1000}) + warn_pattern = re.compile( + rf'Expected response size \(\d*\) larger than ' + rf'download warn size \(1000\) in request {request}' + ) + + yield self._check_log_warnsize(request, warn_pattern, Data.HTML_LARGE) + + @inlineCallbacks + def test_log_received_warnsize(self): + request = Request(url=self.get_url('/no-content-length-header'), meta={'download_warnsize': 10}) + warn_pattern = re.compile( + rf'Received more \(\d*\) bytes than download ' + rf'warn size \(10\) in request {request}' + ) + + yield self._check_log_warnsize(request, warn_pattern, Data.NO_CONTENT_LENGTH) + + def test_max_concurrent_streams(self): + """Send 500 requests at one to check if we can handle + very large number of request. + """ + + def get_deferred(): + return self._check_GET( + Request(self.get_url('/get-data-html-small')), + Data.HTML_SMALL, + 200 + ) + + return self._check_repeat(get_deferred, 500) + + def test_inactive_stream(self): + """Here we send 110 requests considering the MAX_CONCURRENT_STREAMS + by default is 100. After sending the first 100 requests we close the + connection.""" + d_list = [] + + def assert_inactive_stream(failure): + self.assertIsNotNone(failure.check(ResponseFailed)) + self.assertTrue(any( + isinstance(e, InactiveStreamClosed) + for e in failure.value.reasons + )) + + # Send 100 request (we do not check the result) + for _ in range(100): + d = self.make_request(Request(self.get_url('/get-data-html-small'))) + d.addBoth(lambda _: None) + d_list.append(d) + + # Now send 10 extra request and save the response deferred in a list + for _ in range(10): + d = self.make_request(Request(self.get_url('/get-data-html-small'))) + d.addCallback(self.fail) + d.addErrback(assert_inactive_stream) + d_list.append(d) + + # Close the connection now to fire all the extra 10 requests errback + # with InactiveStreamClosed + self.client.transport.loseConnection() + + return DeferredList(d_list, consumeErrors=True, fireOnOneErrback=True) + + def test_invalid_request_type(self): + with self.assertRaises(TypeError): + self.make_request('https://InvalidDataTypePassed.com') + + def test_query_parameters(self): + params = { + 'a': generate_random_string(20), + 'b': generate_random_string(20), + 'c': generate_random_string(20), + 'd': generate_random_string(20) + } + request = Request(self.get_url(f'/query-params?{urlencode(params)}')) + + def assert_query_params(response: Response): + content_encoding = str(response.headers[b'Content-Encoding'], 'utf-8') + data = json.loads(str(response.body, content_encoding)) + self.assertEqual(data, params) + + d = self.make_request(request) + d.addCallback(assert_query_params) + d.addErrback(self.fail) + + return d + + def test_status_codes(self): + def assert_response_status(response: Response, expected_status: int): + self.assertEqual(response.status, expected_status) + + d_list = [] + for status in [200, 404]: + request = Request(self.get_url(f'/status?n={status}')) + d = self.make_request(request) + d.addCallback(assert_response_status, status) + d.addErrback(self.fail) + d_list.append(d) + + return DeferredList(d_list, fireOnOneErrback=True) + + def test_response_has_correct_certificate_ip_address(self): + request = Request(self.get_url('/status?n=200')) + + def assert_metadata(response: Response): + self.assertEqual(response.request, request) + self.assertIsInstance(response.certificate, Certificate) + self.assertIsNotNone(response.certificate.original) + self.assertEqual(response.certificate.getIssuer(), self.client_certificate.getIssuer()) + self.assertTrue(response.certificate.getPublicKey().matches(self.client_certificate.getPublicKey())) + + self.assertIsInstance(response.ip_address, IPv4Address) + self.assertEqual(str(response.ip_address), '127.0.0.1') + + d = self.make_request(request) + d.addCallback(assert_metadata) + d.addErrback(self.fail) + + return d + + def _check_invalid_netloc(self, url): + request = Request(url) + + def assert_invalid_hostname(failure: Failure): + self.assertIsNotNone(failure.check(InvalidHostname)) + error_msg = str(failure.value) + self.assertIn('localhost', error_msg) + self.assertIn('127.0.0.1', error_msg) + self.assertIn(str(request), error_msg) + + d = self.make_request(request) + d.addCallback(self.fail) + d.addErrback(assert_invalid_hostname) + return d + + def test_invalid_hostname(self): + return self._check_invalid_netloc('https://notlocalhost.notlocalhostdomain') + + def test_invalid_host_port(self): + port = self.port_number + 1 + return self._check_invalid_netloc(f'https://127.0.0.1:{port}') + + def test_connection_stays_with_invalid_requests(self): + d_list = [ + self.test_invalid_hostname(), + self.test_invalid_host_port(), + self.test_GET_small_body(), + self.test_POST_small_json() + ] + + return DeferredList(d_list, fireOnOneErrback=True) + + def test_connection_timeout(self): + request = Request(self.get_url('/timeout')) + d = self.make_request(request) + + # Update the timer to 1s to test connection timeout + self.client.setTimeout(1) + + def assert_timeout_error(failure: Failure): + for err in failure.value.reasons: + if isinstance(err, TimeoutError): + self.assertIn(f"Connection was IDLE for more than {H2ClientProtocol.IDLE_TIMEOUT}s", str(err)) + break + else: + self.fail() + + d.addCallback(self.fail) + d.addErrback(assert_timeout_error) + return d + + def test_request_headers_received(self): + request = Request(self.get_url('/request-headers'), headers={ + 'header-1': 'header value 1', + 'header-2': 'header value 2' + }) + d = self.make_request(request) + + def assert_request_headers(response: Response): + self.assertEqual(response.status, 200) + self.assertEqual(response.request, request) + + response_headers = json.loads(str(response.body, 'utf-8')) + self.assertIsInstance(response_headers, dict) + for k, v in request.headers.items(): + k, v = str(k, 'utf-8'), str(v[0], 'utf-8') + self.assertIn(k, response_headers) + self.assertEqual(v, response_headers[k]) + + d.addErrback(self.fail) + d.addCallback(assert_request_headers) + return d diff --git a/tests/upper-constraints.txt b/tests/upper-constraints.txt new file mode 100644 index 000000000..75f337856 --- /dev/null +++ b/tests/upper-constraints.txt @@ -0,0 +1,17 @@ +# Request the latest known version or newer of some dependencies to prevent the +# pip dependency resolver from spending too much time backtracking. +attrs>=20.2.0 +Automat>=0.8.0 +botocore>=1.20.3 +itemadapter>=0.1.1 +itemloaders>=1.0.3 +lxml>=4.6.1 +parsel>=1.5.2 +Pillow>=8.0.1 +pyOpenSSL>=17.5 # mitmproxy 4.0.4 +pytest>=6.2.1 +pytest-twisted>=1.13.1 +service_identity>=17.0.0 +six>=1.14.0 +sybil>=2.0.0 +Twisted>=19.10.0 diff --git a/tox.ini b/tox.ini index 9815f80f7..86ae951b5 100644 --- a/tox.ini +++ b/tox.ini @@ -9,7 +9,6 @@ minversion = 1.7.0 [testenv] deps = - -ctests/constraints.txt -rtests/requirements-py3.txt # mitmproxy does not support PyPy # mitmproxy does not support Windows when running Python < 3.7 @@ -21,7 +20,7 @@ deps = botocore>=1.4.87 Pillow>=4.0.0 # Twisted 21+ causes issues in tests that use skipIf - Twisted<21 + Twisted[http2]>=17.9.0,<21 passenv = S3_TEST_FILE_URI AWS_ACCESS_KEY_ID @@ -32,6 +31,8 @@ passenv = download = true commands = py.test --cov=scrapy --cov-report=xml --cov-report= {posargs:--durations=10 docs scrapy tests} +install_command = + pip install -U -ctests/upper-constraints.txt {opts} {packages} [testenv:typing] basepython = python3 @@ -70,16 +71,16 @@ commands = [pinned] deps = - -ctests/constraints.txt cryptography==2.0 cssselect==0.9.1 + h2==3.2.0 itemadapter==0.1.0 parsel==1.5.0 Protego==0.1.15 pyOpenSSL==16.2.0 queuelib==1.4.2 service_identity==16.0.0 - Twisted==17.9.0 + Twisted[http2]==17.9.0 w3lib==1.17.0 zope.interface==4.1.3 -rtests/requirements-py3.txt @@ -93,12 +94,15 @@ deps = Pillow==4.0.0 setenv = _SCRAPY_PINNED=true +install_command = + pip install -U {opts} {packages} [testenv:pinned] deps = {[pinned]deps} lxml==3.5.0 PyDispatcher==2.0.5 +install_command = {[pinned]install_command} setenv = {[pinned]setenv} @@ -110,6 +114,7 @@ deps = # not need to build lxml from sources in a CI Windows job: lxml==3.8.0 PyDispatcher==2.0.5 +install_command = {[pinned]install_command} setenv = {[pinned]setenv} @@ -126,6 +131,7 @@ commands = [testenv:asyncio-pinned] deps = {[testenv:pinned]deps} commands = {[testenv:asyncio]commands} +install_command = {[pinned]install_command} setenv = {[pinned]setenv} @@ -141,6 +147,7 @@ deps = lxml==4.0.0 PyPyDispatcher==2.1.0 commands = {[testenv:pypy3]commands} +install_command = {[pinned]install_command} setenv = {[pinned]setenv}