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
This commit is contained in:
Aditya 2020-07-22 13:53:46 +05:30
parent e662762e6a
commit 316620b517
3 changed files with 55 additions and 38 deletions

View File

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

View File

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

View File

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