From 92bec38591fccb523c2e643aef70a1f6cd7267ea Mon Sep 17 00:00:00 2001 From: Aditya Date: Wed, 29 Jul 2020 13:43:59 +0530 Subject: [PATCH] feat: MethodNotAllowed405, Content-Length header - add tests to check for Content-Length header - raise MethodNotAllowed405 when remote send 'HTTP/2.0 405 Method Not Allowed' --- scrapy/core/downloader/handlers/http2.py | 15 +++----- scrapy/core/http2/agent.py | 2 +- scrapy/core/http2/protocol.py | 32 +++++++++++++---- scrapy/core/http2/stream.py | 14 ++++++-- tests/test_http2_client_protocol.py | 46 ++++++++++++++++++++++-- 5 files changed, 85 insertions(+), 24 deletions(-) diff --git a/scrapy/core/downloader/handlers/http2.py b/scrapy/core/downloader/handlers/http2.py index 81ea78e69..e9cc5ebbc 100644 --- a/scrapy/core/downloader/handlers/http2.py +++ b/scrapy/core/downloader/handlers/http2.py @@ -3,6 +3,7 @@ from time import time from typing import Optional, Tuple from urllib.parse import urldefrag +from twisted.internet.defer import Deferred from twisted.internet.base import ReactorBase from twisted.internet.error import TimeoutError from twisted.web.client import URI @@ -23,9 +24,6 @@ class H2DownloadHandler: from twisted.internet import reactor self._pool = H2ConnectionPool(reactor, settings) self._context_factory = 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') @classmethod def from_crawler(cls, crawler): @@ -35,8 +33,6 @@ class H2DownloadHandler: agent = ScrapyH2Agent( context_factory=self._context_factory, pool=self._pool, - maxsize=getattr(spider, 'download_maxsize', self._default_maxsize), - warnsize=getattr(spider, 'download_warnsize', self._default_warnsize), crawler=self._crawler ) return agent.download_request(request, spider) @@ -72,15 +68,12 @@ class ScrapyH2Agent: self, context_factory, connect_timeout=10, bind_address: Optional[bytes] = None, pool: H2ConnectionPool = None, - maxsize: int = 0, warnsize: int = 0, crawler=None ) -> None: self._context_factory = context_factory self._connect_timeout = connect_timeout self._bind_address = bind_address self._pool = pool - self._maxsize = maxsize - self._warnsize = warnsize self._crawler = crawler def _get_agent(self, request: Request, timeout: Optional[float]) -> H2Agent: @@ -121,7 +114,7 @@ class ScrapyH2Agent: pool=self._pool ) - def download_request(self, request: Request, spider: Spider): + 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) @@ -135,12 +128,12 @@ class ScrapyH2Agent: return d @staticmethod - def _cb_latency(response: Response, request: Request, start_time: float): + 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): + def _cb_timeout(response: Response, request: Request, timeout: float, timeout_cl) -> Response: if timeout_cl.active(): timeout_cl.cancel() return response diff --git a/scrapy/core/http2/agent.py b/scrapy/core/http2/agent.py index c7a49fd42..e62eef263 100644 --- a/scrapy/core/http2/agent.py +++ b/scrapy/core/http2/agent.py @@ -93,7 +93,7 @@ class H2ConnectionPool: Deferred that fires when all connections have been closed """ for conn in self._connections.values(): - conn.transport.loseConnection() + conn.transport.abortConnection() @implementer(IPolicyForHTTPS) diff --git a/scrapy/core/http2/protocol.py b/scrapy/core/http2/protocol.py index feb034a0c..1ce8b6548 100644 --- a/scrapy/core/http2/protocol.py +++ b/scrapy/core/http2/protocol.py @@ -13,7 +13,7 @@ from h2.events import ( SettingsAcknowledged, StreamEnded, StreamReset, UnknownFrameReceived, WindowUpdated ) -from h2.exceptions import H2Error, ProtocolError +from h2.exceptions import H2Error from twisted.internet.defer import Deferred from twisted.internet.error import TimeoutError from twisted.internet.interfaces import IHandshakeListener, IProtocolNegotiationFactory @@ -39,7 +39,7 @@ class InvalidNegotiatedProtocol(H2Error): self.negotiated_protocol = negotiated_protocol def __str__(self) -> str: - return f'InvalidHostname: Expected h2 as negotiated protocol, received {self.negotiated_protocol}' + return f'InvalidHostname: Expected h2 as negotiated protocol, received {self.negotiated_protocol!r}' class RemoteTerminatedConnection(H2Error): @@ -48,7 +48,15 @@ class RemoteTerminatedConnection(H2Error): self.terminate_event = event def __str__(self) -> str: - return f'RemoteTerminatedConnection: Received GOAWAY frame from {self.remote_ip_address}' + return f'RemoteTerminatedConnection: Received GOAWAY frame from {self.remote_ip_address!r}' + + +class MethodNotAllowed405(H2Error): + def __init__(self, remote_ip_address: Optional[Union[IPv4Address, IPv6Address]]): + self.remote_ip_address = remote_ip_address + + def __str__(self) -> str: + return f"MethodNotAllowed405: Received 'HTTP/2.0 405 Method Not Allowed' from {self.remote_ip_address!r}" @implementer(IHandshakeListener) @@ -217,14 +225,25 @@ class H2ClientProtocol(Protocol, TimeoutMixin): # So, no need to send a GOAWAY frame to the remote self._lose_connection_with_error([InvalidNegotiatedProtocol(negotiated_protocol)]) + 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 ProtocolError as e: + except H2Error as e: # 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. @@ -271,9 +290,10 @@ class H2ClientProtocol(Protocol, TimeoutMixin): for stream in self.streams.values(): if stream.request_sent: - stream.close(StreamCloseReason.CONNECTION_LOST, self._conn_lost_errors, from_protocol=True) + close_reason = StreamCloseReason.CONNECTION_LOST else: - stream.close(StreamCloseReason.INACTIVE, from_protocol=True) + close_reason = StreamCloseReason.INACTIVE + stream.close(close_reason, self._conn_lost_errors, from_protocol=True) self._active_streams -= len(self.streams) self.streams.clear() diff --git a/scrapy/core/http2/stream.py b/scrapy/core/http2/stream.py index 133196797..5bffa67e7 100644 --- a/scrapy/core/http2/stream.py +++ b/scrapy/core/http2/stream.py @@ -189,12 +189,15 @@ class Stream: headers = [ (':method', self._request.method), (':authority', url.netloc), - (':scheme', 'https'), + (':scheme', self._protocol.metadata['uri'].scheme), (':path', path), ] for name, value in self._request.headers.items(): - headers.append((name, value[0])) + headers.append((str(name, 'utf-8'), str(value[0], 'utf-8'))) + + if b'Content-Length' not in self._request.headers.keys(): + headers.append(('Content-Length', str(len(self._request.body)))) return headers @@ -337,6 +340,10 @@ class Stream: 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) @@ -387,7 +394,8 @@ class Stream: self._deferred_response.errback(ResponseFailed(errors)) elif reason is StreamCloseReason.INACTIVE: - self._deferred_response.errback(InactiveStreamClosed(self._request)) + errors.insert(0, InactiveStreamClosed(self._request)) + self._deferred_response.errback(ResponseFailed(errors)) elif reason is StreamCloseReason.INVALID_HOSTNAME: self._deferred_response.errback(InvalidHostname( diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py index 746eef4d6..4926ada14 100644 --- a/tests/test_http2_client_protocol.py +++ b/tests/test_http2_client_protocol.py @@ -11,14 +11,14 @@ 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 URI +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 twisted.internet.error import TimeoutError from scrapy.core.http2.protocol import H2ClientFactory, H2ClientProtocol from scrapy.core.http2.stream import InactiveStreamClosed, InvalidHostname @@ -152,6 +152,19 @@ class QueryParams(LeafResource): 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()) @@ -179,6 +192,7 @@ class Https2ClientProtocolTestCase(TestCase): 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 @@ -488,7 +502,11 @@ class Https2ClientProtocolTestCase(TestCase): d_list = [] def assert_inactive_stream(failure): - self.assertIsNotNone(failure.check(InactiveStreamClosed)) + 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): @@ -616,3 +634,25 @@ class Https2ClientProtocolTestCase(TestCase): 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