mirror of https://github.com/scrapy/scrapy.git
test: query params, certificate & ip_address
- refactor from str.format() to f-strings
This commit is contained in:
parent
50dd9271b4
commit
7b1ad995a4
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
7
setup.py
7
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,
|
||||
)
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue