mirror of https://github.com/scrapy/scrapy.git
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:
parent
e662762e6a
commit
316620b517
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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'],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue