From 316620b517207b1082dd9c5b4ebfc7fbe745e3bb Mon Sep 17 00:00:00 2001 From: Aditya Date: Wed, 22 Jul 2020 13:53:46 +0530 Subject: [PATCH] chore: pass spider as argument for request method - download_maxsize and download_warnsize can now be extracted from the spider directly and passed to the stream - remove `partial` flag from the response as per RFC 7540 - Section 8.1.2.6 --- scrapy/core/http2/protocol.py | 12 +++++--- scrapy/core/http2/stream.py | 35 ++++++++++------------ tests/test_http2_client_protocol.py | 46 +++++++++++++++++++---------- 3 files changed, 55 insertions(+), 38 deletions(-) diff --git a/scrapy/core/http2/protocol.py b/scrapy/core/http2/protocol.py index 041908116..feb034a0c 100644 --- a/scrapy/core/http2/protocol.py +++ b/scrapy/core/http2/protocol.py @@ -28,6 +28,7 @@ from scrapy.core.http2.stream import Stream, StreamCloseReason from scrapy.core.http2.types import H2ConnectionMetadataDict from scrapy.http import Request from scrapy.settings import Settings +from scrapy.spiders import Spider logger = logging.getLogger(__name__) @@ -71,7 +72,7 @@ class H2ClientProtocol(Protocol, TimeoutMixin): # ID of the next request stream # Following the convention - 'Streams initiated by a client MUST - # use odd-numbered stream identifiers' (RFC 7540) + # 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 @@ -136,6 +137,7 @@ class H2ClientProtocol(Protocol, TimeoutMixin): self._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 @@ -145,13 +147,15 @@ class H2ClientProtocol(Protocol, TimeoutMixin): self._send_pending_requests() return stream - def _new_stream(self, request: Request) -> 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 @@ -166,11 +170,11 @@ class H2ClientProtocol(Protocol, TimeoutMixin): data = self.conn.data_to_send() self.transport.write(data) - def request(self, request: Request) -> Deferred: + 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) + stream = self._new_stream(request, spider) d = stream.get_response() # Add the stream to the request pool diff --git a/scrapy/core/http2/stream.py b/scrapy/core/http2/stream.py index 15f081cab..133196797 100644 --- a/scrapy/core/http2/stream.py +++ b/scrapy/core/http2/stream.py @@ -83,7 +83,9 @@ class Stream: self, stream_id: int, request: Request, - protocol: "H2ClientProtocol" + protocol: "H2ClientProtocol", + download_maxsize: int = 0, + download_warnsize: int = 0 ) -> None: """ Arguments: @@ -95,14 +97,8 @@ class Stream: self._request: Request = request self._protocol: "H2ClientProtocol" = protocol - self._download_maxsize = self._request.meta.get( - 'download_maxsize', - self._protocol.metadata['default_download_maxsize'] - ) - self._download_warnsize = self._request.meta.get( - 'download_warnsize', - self._protocol.metadata['default_download_warnsize'] - ) + self._download_maxsize = self._request.meta.get('download_maxsize', download_maxsize) + self._download_warnsize = self._request.meta.get('download_warnsize', download_warnsize) self.request_start_time = None @@ -346,26 +342,28 @@ class Stream: self.stream_closed_server = True - flags = None - if b'Content-Length' not in self._response['headers']: - # Missing Content-Length - {twisted.web.http.PotentialDataLoss} - flags = ['partial'] + # 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', -1)) + expected_size = int(self._response['headers'].get( + b'Content-Length', + self._response['flow_controlled_size']) + ) error_msg = ( - f'Cancelling download of {self._request.url}: expected response ' - f'size ({expected_size}) larger than download max size ({self._download_maxsize}).' + 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(flags) + self._fire_response_deferred() # Stream was abruptly ended here elif reason is StreamCloseReason.CANCELLED: @@ -398,7 +396,7 @@ class Stream: f'{self._protocol.metadata["ip_address"]}:{self._protocol.metadata["uri"].port}' )) - def _fire_response_deferred(self, flags: Optional[List[str]] = None) -> None: + def _fire_response_deferred(self) -> None: """Builds response from the self._response dict and fires the response deferred callback with the generated response instance""" @@ -416,7 +414,6 @@ class Stream: headers=self._response['headers'], body=body, request=self._request, - flags=flags, certificate=self._protocol.metadata['certificate'], ip_address=self._protocol.metadata['ip_address'], ) diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py index 2833801e7..d0386f7f8 100644 --- a/tests/test_http2_client_protocol.py +++ b/tests/test_http2_client_protocol.py @@ -23,6 +23,7 @@ from scrapy.core.http2.protocol import H2ClientFactory 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 @@ -41,6 +42,14 @@ def make_html_body(val): return bytes(response, 'utf-8') +class DummySpider(Spider): + name = 'dummy' + start_urls = [] + + def parse(self, response): + print(response) + + class Data: SMALL_SIZE = 1024 # 1 KB LARGE_SIZE = 1024 ** 2 # 1 MB @@ -210,6 +219,9 @@ class Https2ClientProtocolTestCase(TestCase): 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 = [] @@ -233,7 +245,7 @@ class Https2ClientProtocolTestCase(TestCase): content_length = int(response.headers.get('Content-Length')) self.assertEqual(len(response.body), content_length) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(check_response) d.addErrback(self.fail) return d @@ -273,7 +285,7 @@ class Https2ClientProtocolTestCase(TestCase): expected_extra_data, expected_status: int ): - d = self.client.request(request) + d = self.make_request(request) def assert_response(response: Response): self.assertEqual(response.status, expected_status) @@ -355,7 +367,7 @@ class Https2ClientProtocolTestCase(TestCase): self.assertEqual(response.status, 499) self.assertEqual(response.request, request) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(assert_response) d.addErrback(self.fail) d.cancel() @@ -367,8 +379,13 @@ class Https2ClientProtocolTestCase(TestCase): 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.client.request(request) + d = self.make_request(request) d.addCallback(self.fail) d.addErrback(assert_cancelled_error) return d @@ -385,7 +402,7 @@ class Https2ClientProtocolTestCase(TestCase): for error in failure.value.reasons )) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(self.fail) d.addErrback(assert_failure) return d @@ -397,10 +414,9 @@ class Https2ClientProtocolTestCase(TestCase): self.assertEqual(response.status, 200) self.assertEqual(response.body, Data.NO_CONTENT_LENGTH) self.assertEqual(response.request, request) - self.assertIn('partial', response.flags) self.assertNotIn('Content-Length', response.headers) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(assert_content_length) d.addErrback(self.fail) return d @@ -413,7 +429,7 @@ class Https2ClientProtocolTestCase(TestCase): expected_body ): with self.assertLogs('scrapy.core.http2.stream', level='WARNING') as cm: - response = yield self.client.request(request) + response = yield self.make_request(request) self.assertEqual(response.status, 200) self.assertEqual(response.request, request) self.assertEqual(response.body, expected_body) @@ -469,13 +485,13 @@ class Https2ClientProtocolTestCase(TestCase): # Send 100 request (we do not check the result) for _ in range(100): - d = self.client.request(Request(self.get_url('/get-data-html-small'))) + 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.client.request(Request(self.get_url('/get-data-html-small'))) + 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) @@ -488,7 +504,7 @@ class Https2ClientProtocolTestCase(TestCase): def test_invalid_request_type(self): with self.assertRaises(TypeError): - self.client.request('https://InvalidDataTypePassed.com') + self.make_request('https://InvalidDataTypePassed.com') def test_query_parameters(self): params = { @@ -504,7 +520,7 @@ class Https2ClientProtocolTestCase(TestCase): data = json.loads(str(response.body, content_encoding)) self.assertEqual(data, params) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(assert_query_params) d.addErrback(self.fail) @@ -517,7 +533,7 @@ class Https2ClientProtocolTestCase(TestCase): d_list = [] for status in [200, 404]: request = Request(self.get_url(f'/status?n={status}')) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(assert_response_status, status) d.addErrback(self.fail) d_list.append(d) @@ -537,7 +553,7 @@ class Https2ClientProtocolTestCase(TestCase): self.assertIsInstance(response.ip_address, IPv4Address) self.assertEqual(str(response.ip_address), '127.0.0.1') - d = self.client.request(request) + d = self.make_request(request) d.addCallback(assert_metadata) d.addErrback(self.fail) @@ -553,7 +569,7 @@ class Https2ClientProtocolTestCase(TestCase): self.assertIn('127.0.0.1', error_msg) self.assertIn(str(request), error_msg) - d = self.client.request(request) + d = self.make_request(request) d.addCallback(self.fail) d.addErrback(assert_invalid_hostname) return d