diff --git a/scrapy/core/spidermw.py b/scrapy/core/spidermw.py index 763e0cdf6..961606f29 100644 --- a/scrapy/core/spidermw.py +++ b/scrapy/core/spidermw.py @@ -3,6 +3,7 @@ Spider Middleware manager See documentation in docs/topics/spider-middleware.rst """ +import inspect from itertools import islice from twisted.python.failure import Failure @@ -11,11 +12,11 @@ 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.python import MutableChain +from scrapy.utils.python import MutableAsyncChain, MutableChain def _isiterable(possible_iterator): - return hasattr(possible_iterator, '__iter__') + return hasattr(possible_iterator, '__iter__') or hasattr(possible_iterator, '__aiter__') def _fname(f): @@ -58,15 +59,31 @@ class SpiderMiddlewareManager(MiddlewareManager): return scrape_func(response, request, spider) def _evaluate_iterable(iterable, exception_processor_index, recover_to): - try: - for r in iterable: - yield r - except Exception as ex: + def _process_exception(ex): exception_result = process_spider_exception(Failure(ex), exception_processor_index) if isinstance(exception_result, Failure): raise recover_to.extend(exception_result) + def _evaluate_normal_iterable(iterable): + try: + for r in iterable: + yield r + except Exception as ex: + _process_exception(ex) + + async def _evaluate_async_iterable(iterable): + try: + async for r in iterable: + yield r + except Exception as ex: + _process_exception(ex) + + if inspect.isasyncgen(iterable): + return _evaluate_async_iterable(iterable) + else: + return _evaluate_normal_iterable(iterable) + def process_spider_exception(_failure, start_index=0): exception = _failure.value # don't handle _InvalidOutput exception @@ -92,7 +109,11 @@ class SpiderMiddlewareManager(MiddlewareManager): def process_spider_output(result, start_index=0): # items in this iterable do not need to go through the process_spider_output # chain, they went through it already from the process_spider_exception method - recovered = MutableChain() + if inspect.isasyncgen(result): + iter_class = MutableAsyncChain + else: + iter_class = MutableChain + recovered = iter_class() method_list = islice(self.methods['process_spider_output'], start_index, None) for method_index, method in enumerate(method_list, start=start_index): @@ -113,12 +134,16 @@ class SpiderMiddlewareManager(MiddlewareManager): f"iterable, got {type(result)}") raise _InvalidOutput(msg) - return MutableChain(result, recovered) + return iter_class(result, recovered) def process_callback_output(result): - recovered = MutableChain() + if inspect.isasyncgen(result): + iter_class = MutableAsyncChain + else: + iter_class = MutableChain + recovered = iter_class() result = _evaluate_iterable(result, 0, recovered) - return MutableChain(process_spider_output(result), recovered) + return iter_class(process_spider_output(result), recovered) dfd = mustbe_deferred(process_spider_input, response) dfd.addCallbacks(callback=process_callback_output, errback=process_spider_exception) diff --git a/scrapy/utils/defer.py b/scrapy/utils/defer.py index bd5f9c8fc..554edc38c 100644 --- a/scrapy/utils/defer.py +++ b/scrapy/utils/defer.py @@ -290,3 +290,22 @@ def maybeDeferred_coro(f, *args, **kw): return defer.fail(result) else: return defer.succeed(result) + + +def deferred_to_future(d): + """ Wraps a Deferred into a Future. Requires the asyncio reactor. + """ + return d.asFuture(asyncio.get_event_loop()) + + +def maybe_deferred_to_future(d): + """ Converts a Deferred to something that can be awaited in a callback or other user coroutine. + + If the asyncio reactor is installed, coroutines are wrapped into Futures, and only Futures can be + awaited inside them. Otherwise, coroutines are wrapped into Deferreds and Deferreds can be awaited + directly inside them. + """ + if not is_asyncio_reactor_installed(): + return d + else: + return deferred_to_future(d) diff --git a/scrapy/utils/spider.py b/scrapy/utils/spider.py index 59fc9202f..d0fd1757d 100644 --- a/scrapy/utils/spider.py +++ b/scrapy/utils/spider.py @@ -4,7 +4,6 @@ import logging from scrapy.spiders import Spider from scrapy.utils.defer import deferred_from_coro from scrapy.utils.misc import arg_to_iter -from scrapy.utils.asyncgen import collect_asyncgen logger = logging.getLogger(__name__) @@ -12,14 +11,13 @@ logger = logging.getLogger(__name__) def iterate_spider_output(result): if inspect.isasyncgen(result): - d = deferred_from_coro(collect_asyncgen(result)) - d.addCallback(iterate_spider_output) - return d + return result elif inspect.iscoroutine(result): d = deferred_from_coro(result) d.addCallback(iterate_spider_output) return d - return arg_to_iter(result) + else: + return arg_to_iter(deferred_from_coro(result)) def iter_spider_classes(module): diff --git a/scrapy/utils/test.py b/scrapy/utils/test.py index 24c38283a..d8fc25094 100644 --- a/scrapy/utils/test.py +++ b/scrapy/utils/test.py @@ -110,3 +110,10 @@ def mock_google_cloud_storage(): bucket_mock.blob.return_value = blob_mock return (client_mock, bucket_mock, blob_mock) + + +def get_web_client_agent_req(url): + from twisted.internet import reactor + from twisted.web.client import Agent # imports twisted.internet.reactor + agent = Agent(reactor) + return agent.request(b'GET', url.encode('utf-8')) diff --git a/tests/spiders.py b/tests/spiders.py index 106392ea6..3e0ec001b 100644 --- a/tests/spiders.py +++ b/tests/spiders.py @@ -14,7 +14,8 @@ from scrapy.item import Item from scrapy.linkextractors import LinkExtractor from scrapy.spiders import Spider from scrapy.spiders.crawl import CrawlSpider, Rule -from scrapy.utils.test import get_from_asyncio_queue +from scrapy.utils.defer import deferred_to_future, maybe_deferred_to_future +from scrapy.utils.test import get_from_asyncio_queue, get_web_client_agent_req class MockServerSpider(Spider): @@ -148,6 +149,41 @@ class AsyncDefAsyncioReqsReturnSpider(SimpleSpider): return reqs +class AsyncDefAsyncioGenExcSpider(SimpleSpider): + name = 'asyncdef_asyncio_gen_exc' + + async def parse(self, response): + for i in range(10): + await asyncio.sleep(0.1) + yield {'foo': i} + if i > 5: + raise ValueError("Stopping the processing") + + +class AsyncDefDeferredDirectSpider(SimpleSpider): + name = 'asyncdef_deferred_direct' + + async def parse(self, response): + resp = await get_web_client_agent_req(self.mockserver.url("/status?n=200")) + yield {'code': resp.code} + + +class AsyncDefDeferredWrappedSpider(SimpleSpider): + name = 'asyncdef_deferred_wrapped' + + async def parse(self, response): + resp = await deferred_to_future(get_web_client_agent_req(self.mockserver.url("/status?n=200"))) + yield {'code': resp.code} + + +class AsyncDefDeferredMaybeWrappedSpider(SimpleSpider): + name = 'asyncdef_deferred_wrapped' + + async def parse(self, response): + resp = await maybe_deferred_to_future(get_web_client_agent_req(self.mockserver.url("/status?n=200"))) + yield {'code': resp.code} + + class AsyncDefAsyncioGenSpider(SimpleSpider): name = 'asyncdef_asyncio_gen' diff --git a/tests/test_crawl.py b/tests/test_crawl.py index 1083c1678..cda52f0d4 100644 --- a/tests/test_crawl.py +++ b/tests/test_crawl.py @@ -20,12 +20,16 @@ from scrapy.utils.python import to_unicode from tests.mockserver import MockServer from tests.spiders import ( AsyncDefAsyncioGenComplexSpider, + AsyncDefAsyncioGenExcSpider, AsyncDefAsyncioGenLoopSpider, AsyncDefAsyncioGenSpider, AsyncDefAsyncioReqsReturnSpider, AsyncDefAsyncioReturnSingleElementSpider, AsyncDefAsyncioReturnSpider, AsyncDefAsyncioSpider, + AsyncDefDeferredDirectSpider, + AsyncDefDeferredMaybeWrappedSpider, + AsyncDefDeferredWrappedSpider, AsyncDefSpider, BrokenStartRequestsSpider, BytesReceivedCallbackSpider, @@ -430,6 +434,18 @@ class CrawlSpiderTestCase(TestCase): for i in range(10): self.assertIn({'foo': i}, items) + @mark.only_asyncio() + @defer.inlineCallbacks + def test_async_def_asyncgen_parse_exc(self): + log, items, stats = yield self._run_spider(AsyncDefAsyncioGenExcSpider) + log = str(log) + self.assertIn("Spider error processing", log) + self.assertIn("ValueError", log) + itemcount = stats.get_value('item_scraped_count') + self.assertEqual(itemcount, 7) + for i in range(7): + self.assertIn({'foo': i}, items) + @mark.only_asyncio() @defer.inlineCallbacks def test_async_def_asyncgen_parse_complex(self): @@ -449,6 +465,23 @@ class CrawlSpiderTestCase(TestCase): for req_id in range(3): self.assertIn(f"Got response 200, req_id {req_id}", str(log)) + @mark.only_not_asyncio() + @defer.inlineCallbacks + def test_async_def_deferred_direct(self): + _, items, _ = yield self._run_spider(AsyncDefDeferredDirectSpider) + self.assertEqual(items, [{'code': 200}]) + + @mark.only_asyncio() + @defer.inlineCallbacks + def test_async_def_deferred_wrapped(self): + log, items, _ = yield self._run_spider(AsyncDefDeferredWrappedSpider) + self.assertEqual(items, [{'code': 200}]) + + @defer.inlineCallbacks + def test_async_def_deferred_maybe_wrapped(self): + _, items, _ = yield self._run_spider(AsyncDefDeferredMaybeWrappedSpider) + self.assertEqual(items, [{'code': 200}]) + @defer.inlineCallbacks def test_response_ssl_certificate_none(self): crawler = self.runner.create_crawler(SingleRequestSpider) diff --git a/tests/test_spidermiddleware_output_chain.py b/tests/test_spidermiddleware_output_chain.py index 029bf8bd6..4d1a7fcb0 100644 --- a/tests/test_spidermiddleware_output_chain.py +++ b/tests/test_spidermiddleware_output_chain.py @@ -43,6 +43,23 @@ class RecoverySpider(Spider): raise TabError() +class RecoveryAsyncGenSpider(RecoverySpider): + name = 'RecoveryAsyncGenSpider' + + async def parse(self, response): + for r in super().parse(response): + yield r + + +class RecoveryMiddleware: + def process_spider_exception(self, response, exception, spider): + spider.logger.info('Middleware: %s exception caught', exception.__class__.__name__) + return [ + {'from': 'process_spider_exception'}, + Request(response.url, meta={'dont_fail': True}, dont_filter=True), + ] + + # ================================================================================ # (1) exceptions from a spider middleware's process_spider_input method class FailProcessSpiderInputMiddleware: @@ -307,6 +324,16 @@ class TestSpiderMiddleware(TestCase): self.assertEqual(str(log).count("Middleware: TabError exception caught"), 1) self.assertIn("'item_scraped_count': 3", str(log)) + @defer.inlineCallbacks + def test_recovery_asyncgen(self): + """ + Same as test_recovery but with an async callback. + """ + log = yield self.crawl_log(RecoveryAsyncGenSpider) + self.assertIn("Middleware: TabError exception caught", str(log)) + self.assertEqual(str(log).count("Middleware: TabError exception caught"), 1) + self.assertIn("'item_scraped_count': 3", str(log)) + @defer.inlineCallbacks def test_process_spider_input_without_errback(self): """