From 7b1ad995a4996babf9a019815dd7256a1cbfa044 Mon Sep 17 00:00:00 2001 From: Aditya Date: Wed, 1 Jul 2020 10:45:36 +0530 Subject: [PATCH] test: query params, certificate & ip_address - refactor from str.format() to f-strings --- scrapy/core/http2/protocol.py | 12 ++-- scrapy/core/http2/stream.py | 46 +++++--------- scrapy/core/http2/types.py | 6 +- setup.py | 7 ++- tests/test_http2_client_protocol.py | 96 +++++++++++++++++++++++------ 5 files changed, 108 insertions(+), 59 deletions(-) diff --git a/scrapy/core/http2/protocol.py b/scrapy/core/http2/protocol.py index 4de80c05e..3438c99f0 100644 --- a/scrapy/core/http2/protocol.py +++ b/scrapy/core/http2/protocol.py @@ -122,6 +122,9 @@ class H2ClientProtocol(Protocol): self.transport.write(data) def request(self, request: Request): + if not isinstance(request, Request): + raise TypeError(f'Expected type scrapy.http.Request but received {request.__class__.__name__}') + stream = self._new_stream(request) d = stream.get_response() @@ -134,9 +137,7 @@ class H2ClientProtocol(Protocol): sending some data now: we should open with the connection preamble. """ self.destination = self.transport.getPeer() - logger.info('Connection made to {}'.format(self.destination)) - - self._metadata['certificate'] = Certificate(self.transport.getPeerCertificate()) + logger.info(f'Connection made to {self.destination}') self._metadata['ip_address'] = ipaddress.ip_address(self.destination.host) self.conn.initiate_connection() @@ -200,7 +201,7 @@ class H2ClientProtocol(Protocol): elif isinstance(event, SettingsAcknowledged): self.settings_acknowledged(event) else: - logger.debug("Received unhandled event {}".format(event)) + logger.debug(f'Received unhandled event {event}') # Event handler functions starts here def data_received(self, event: DataReceived): @@ -214,6 +215,9 @@ class H2ClientProtocol(Protocol): # 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): self.streams[event.stream_id].close(StreamCloseReason.ENDED) diff --git a/scrapy/core/http2/stream.py b/scrapy/core/http2/stream.py index 8d0c6d94d..f4a90a753 100644 --- a/scrapy/core/http2/stream.py +++ b/scrapy/core/http2/stream.py @@ -131,7 +131,7 @@ class Stream: self._deferred_response = Deferred(_cancel) def __str__(self): - return "Stream(id={})".format(repr(self.stream_id)) + return f'Stream(id={self.stream_id})' __repr__ = __str__ @@ -167,13 +167,15 @@ class Stream: url = urlparse(self._request.url) # Make sure pseudo-headers comes before all the other headers + path = url.path + if url.query: + path += '?' + url.query + headers = [ (':method', self._request.method), (':authority', url.netloc), - - # TODO: Check if scheme can be 'http' for HTTP/2 ? (':scheme', 'https'), - (':path', url.path), + (':path', path), ] for name, value in self._request.headers.items(): @@ -253,14 +255,9 @@ class Stream: if self._log_warnsize: self._reached_warnsize = True - warning_msg = 'Received more ({bytes}) bytes than download ' \ - + 'warn size ({warnsize}) in request {request}' - warning_args = { - 'bytes': self._response['flow_controlled_size'], - 'warnsize': self._download_warnsize, - 'request': self._request - } - logger.warning(warning_msg.format(**warning_args)) + 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._conn.acknowledge_received_data( @@ -280,14 +277,9 @@ class Stream: if self._log_warnsize: self._reached_warnsize = True - warning_msg = 'Expected response size ({size}) larger than ' \ - + 'download warn size ({warnsize}) in request {request}' - warning_args = { - 'size': expected_size, - 'warnsize': self._download_warnsize, - 'request': self._request - } - logger.warning(warning_msg.format(**warning_args)) + 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.RESET): """Close this stream by sending a RST_FRAME to the remote peer""" @@ -332,16 +324,10 @@ class Stream: # having Content-Length) if reason is StreamCloseReason.MAXSIZE_EXCEEDED: expected_size = int(self._response['headers'].get(b'Content-Length', -1)) - error_msg = ("Cancelling download of {url}: expected response " - "size ({size}) larger than download max size ({maxsize}).") - error_args = { - 'url': self._request.url, - 'size': expected_size, - 'maxsize': self._download_maxsize - } - - logger.error(error_msg, error_args) - self._deferred_response.errback(CancelledError(error_msg.format(**error_args))) + error_msg = f'Cancelling download of {self._request.url}: expected 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) diff --git a/scrapy/core/http2/types.py b/scrapy/core/http2/types.py index f28bf9472..c0961cd3a 100644 --- a/scrapy/core/http2/types.py +++ b/scrapy/core/http2/types.py @@ -1,6 +1,6 @@ from io import BytesIO from ipaddress import IPv4Address, IPv6Address -from typing import Union +from typing import Union, Optional from twisted.internet.ssl import Certificate # for python < 3.8 -- typing.TypedDict is undefined @@ -13,8 +13,8 @@ class H2ConnectionMetadataDict(TypedDict): """Some meta data of this connection initialized when connection is successfully made """ - certificate: Union[None, Certificate] - ip_address: Union[None, IPv4Address, IPv6Address] + certificate: Optional[Certificate] + ip_address: Optional[Union[IPv4Address, IPv6Address]] class H2ResponseDict(TypedDict): diff --git a/setup.py b/setup.py index 8e50733e6..47c5906e4 100644 --- a/setup.py +++ b/setup.py @@ -1,8 +1,8 @@ from os.path import dirname, join - from pkg_resources import parse_version from setuptools import setup, find_packages, __version__ as setuptools_version + with open(join(dirname(__file__), 'scrapy/VERSION'), 'rb') as f: version = f.read().decode('ascii').strip() @@ -25,11 +25,12 @@ if has_environment_marker_platform_impl_support(): 'PyPyDispatcher>=2.1.0', ] + setup( name='Scrapy', version=version, url='https://scrapy.org', - project_urls={ + project_urls = { 'Documentation': 'https://docs.scrapy.org/', 'Source': 'https://github.com/scrapy/scrapy', 'Tracker': 'https://github.com/scrapy/scrapy/issues', @@ -83,4 +84,4 @@ setup( 'typing_extensions>=3.7' ], extras_require=extras_require, -) +) \ No newline at end of file diff --git a/tests/test_http2_client_protocol.py b/tests/test_http2_client_protocol.py index 79c129d11..c6a9bd5fb 100644 --- a/tests/test_http2_client_protocol.py +++ b/tests/test_http2_client_protocol.py @@ -4,13 +4,15 @@ import random import re import shutil import string +from ipaddress import IPv4Address +from urllib.parse import urlencode from h2.exceptions import InvalidBodyLengthError from twisted.internet import reactor from twisted.internet.defer import inlineCallbacks, DeferredList, CancelledError from twisted.internet.endpoints import SSL4ClientEndpoint, SSL4ServerEndpoint, TCP4ServerEndpoint from twisted.internet.protocol import Factory -from twisted.internet.ssl import optionsForClientTLS, PrivateCertificate +from twisted.internet.ssl import optionsForClientTLS, PrivateCertificate, Certificate from twisted.python.failure import Failure from twisted.trial.unittest import TestCase from twisted.web.http import Request as TxRequest @@ -21,7 +23,7 @@ from scrapy.core.http2.protocol import H2ClientProtocol from scrapy.core.http2.stream import InactiveStreamClosed from scrapy.http import Request, Response, JsonRequest from scrapy.utils.python import to_bytes, to_unicode -from tests.mockserver import ssl_context_factory, LeafResource +from tests.mockserver import ssl_context_factory, LeafResource, Status def generate_random_string(size): @@ -32,10 +34,10 @@ def generate_random_string(size): def make_html_body(val): - response = ''' + response = f'''

Hello from HTTP2

-

{}

-'''.format(val) +

{val}

+''' return to_bytes(response) @@ -122,6 +124,17 @@ class NoContentLengthHeader(LeafResource): request.finish() +class QueryParams(LeafResource): + def render_GET(self, request: TxRequest): + request.setHeader('Content-Type', 'application/json') + + query_params = {} + for k, v in request.args.items(): + query_params[to_unicode(k)] = to_unicode(v[0]) + + return to_bytes(json.dumps(query_params)) + + def get_client_certificate(key_file, certificate_file): with open(key_file, 'r') as key, open(certificate_file, 'r') as certificate: pem = ''.join(key.readlines()) + ''.join(certificate.readlines()) @@ -146,6 +159,8 @@ class Https2ClientProtocolTestCase(TestCase): r.putChild(b'dataloss', Dataloss()) r.putChild(b'no-content-length-header', NoContentLengthHeader()) + r.putChild(b'status', Status()) + r.putChild(b'query-params', QueryParams()) return r @inlineCallbacks @@ -189,7 +204,7 @@ class Https2ClientProtocolTestCase(TestCase): :return: Complete url """ assert len(path) > 0 and (path[0] == '/' or path[0] == '&') - return "{}://{}:{}{}".format(self.scheme, self.hostname, self.port_number, path) + return f'{self.scheme}://{self.hostname}:{self.port_number}{path}' @staticmethod def _check_repeat(get_deferred, count): @@ -210,7 +225,6 @@ class Https2ClientProtocolTestCase(TestCase): self.assertEqual(response.status, expected_status) self.assertEqual(response.body, expected_body) self.assertEqual(response.request, request) - self.assertEqual(response.url, request.url) content_length = int(response.headers.get('Content-Length')) self.assertEqual(len(response.body), content_length) @@ -260,7 +274,6 @@ class Https2ClientProtocolTestCase(TestCase): def assert_response(response: Response): self.assertEqual(response.status, expected_status) self.assertEqual(response.request, request) - self.assertEqual(response.url, request.url) content_length = int(response.headers.get('Content-Length')) self.assertEqual(len(response.body), content_length) @@ -336,7 +349,6 @@ class Https2ClientProtocolTestCase(TestCase): def assert_response(response: Response): self.assertEqual(response.status, 499) self.assertEqual(response.request, request) - self.assertEqual(response.url, request.url) d = self.client.request(request) d.addCallback(assert_response) @@ -356,10 +368,6 @@ class Https2ClientProtocolTestCase(TestCase): d.addErrback(assert_cancelled_error) return d - # TODO: Test in multiple requests if one request fails due to dataloss - # remaining request do not fail (change expected behaviour) - # Can be done only when hyper-h2 don't terminate connection over - # InvalidBodyLengthError check def test_received_dataloss_response(self): """In case when value of Header Content-Length != len(Received Data) ProtocolError is raised""" @@ -384,7 +392,6 @@ class Https2ClientProtocolTestCase(TestCase): self.assertEqual(response.status, 200) self.assertEqual(response.body, Data.NO_CONTENT_LENGTH) self.assertEqual(response.request, request) - self.assertEqual(response.url, request.url) self.assertIn('partial', response.flags) self.assertNotIn('Content-Length', response.headers) @@ -404,7 +411,6 @@ class Https2ClientProtocolTestCase(TestCase): response = yield self.client.request(request) self.assertEqual(response.status, 200) self.assertEqual(response.request, request) - self.assertEqual(response.url, request.url) self.assertEqual(response.body, expected_body) # Check the warning is raised only once for this request @@ -417,8 +423,8 @@ class Https2ClientProtocolTestCase(TestCase): def test_log_expected_warnsize(self): request = Request(url=self.get_url('/get-data-html-large'), meta={'download_warnsize': 1000}) warn_pattern = re.compile( - r'Expected response size \(\d*\) larger than ' - r'download warn size \(1000\) in request {}'.format(request) + 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) @@ -427,8 +433,8 @@ class Https2ClientProtocolTestCase(TestCase): def test_log_received_warnsize(self): request = Request(url=self.get_url('/no-content-length-header'), meta={'download_warnsize': 10}) warn_pattern = re.compile( - r'Received more \(\d*\) bytes than download ' - r'warn size \(10\) in request {}'.format(request) + 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) @@ -474,3 +480,55 @@ class Https2ClientProtocolTestCase(TestCase): self.client.transport.abortConnection() return DeferredList(d_list, consumeErrors=True, fireOnOneErrback=True) + + def test_invalid_request_type(self): + with self.assertRaises(TypeError): + self.client.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): + data = json.loads(to_unicode(response.body)) + self.assertEqual(data, params) + + d = self.client.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.client.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.assertIsInstance(response.ip_address, IPv4Address) + self.assertEqual(str(response.ip_address), '127.0.0.1') + + d = self.client.request(request) + d.addCallback(assert_metadata) + d.addErrback(self.fail) + + return d