test: query params, certificate & ip_address

- refactor from str.format() to f-strings
This commit is contained in:
Aditya 2020-07-01 10:45:36 +05:30
parent 50dd9271b4
commit 7b1ad995a4
5 changed files with 108 additions and 59 deletions

View File

@ -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)

View File

@ -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)

View File

@ -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):

View File

@ -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,
)
)

View File

@ -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 = '''<html>
response = f'''<html>
<h1>Hello from HTTP2<h1>
<p>{}</p>
</html>'''.format(val)
<p>{val}</p>
</html>'''
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