diff --git a/scrapy/core/spidermw.py b/scrapy/core/spidermw.py index 5a99b96be..dd0cf675f 100644 --- a/scrapy/core/spidermw.py +++ b/scrapy/core/spidermw.py @@ -10,7 +10,7 @@ from twisted.python.failure import Failure from scrapy.exceptions import _InvalidOutput from scrapy.middleware import MiddlewareManager from scrapy.utils.conf import build_component_list -from scrapy.utils.defer import mustbe_deferred +from scrapy.utils.defer import mustbe_deferred, deferred_f_from_coro_f from scrapy.utils.python import MutableChain @@ -38,7 +38,7 @@ class SpiderMiddlewareManager(MiddlewareManager): if hasattr(mw, 'process_spider_input'): self.methods['process_spider_input'].append(mw.process_spider_input) if hasattr(mw, 'process_start_requests'): - self.methods['process_start_requests'].appendleft(mw.process_start_requests) + self.methods['process_start_requests'].appendleft(deferred_f_from_coro_f(mw.process_start_requests)) process_spider_output = getattr(mw, 'process_spider_output', None) self.methods['process_spider_output'].appendleft(process_spider_output) process_spider_exception = getattr(mw, 'process_spider_exception', None) diff --git a/tests/test_spidermiddleware.py b/tests/test_spidermiddleware.py index 7adf9a84b..6f931e8f8 100644 --- a/tests/test_spidermiddleware.py +++ b/tests/test_spidermiddleware.py @@ -1,3 +1,4 @@ +import collections import inspect import sys from unittest import mock @@ -115,6 +116,11 @@ class ProcessStartRequestsSimpleMiddleware: yield r +class ProcessStartRequestsAsyncDefMiddleware: + async def process_start_requests(self, start_requests, spider): + return start_requests + + class ProcessStartRequestsSimple(TestCase): """ process_start_requests tests for simple start_requests""" @@ -139,7 +145,7 @@ class ProcessStartRequestsSimple(TestCase): @defer.inlineCallbacks def _test_simple_base(self, *mw_classes): processed_start_requests = yield self._get_processed_start_requests(*mw_classes) - self.assertTrue(inspect.isgenerator(processed_start_requests)) + self.assertIsInstance(processed_start_requests, collections.abc.Iterable) start_requests_list = list(processed_start_requests) self.assertEqual(len(start_requests_list), 3) self.assertIsInstance(start_requests_list[0], Request) @@ -158,6 +164,11 @@ class ProcessStartRequestsSimple(TestCase): """ Simple mw """ yield self._test_simple_base(ProcessStartRequestsSimpleMiddleware) + @defer.inlineCallbacks + def test_asyncdef(self): + """ Async def mw """ + yield self._test_simple_base(ProcessStartRequestsAsyncDefMiddleware) + @mark.skipif(sys.version_info < (3, 6), reason="Async generators require Python 3.6 or higher") @defer.inlineCallbacks def test_asyncgen(self): @@ -208,6 +219,11 @@ class ProcessStartRequestsAsyncGen(ProcessStartRequestsSimple): self.assertTrue(inspect.isgenerator(processed_start_requests)) self.assertAsyncGeneratorNotIterable(processed_start_requests) + @defer.inlineCallbacks + def test_asyncdef(self): + """ Async def mw """ + yield self._test_asyncgen_base(ProcessStartRequestsAsyncDefMiddleware) + @defer.inlineCallbacks def test_simple_asyncgen(self): """ Simple mw -> asyncgen mw; cannot work """