diff --git a/scrapy/crawler.py b/scrapy/crawler.py index 20990ea41..c982611fb 100644 --- a/scrapy/crawler.py +++ b/scrapy/crawler.py @@ -1,3 +1,4 @@ +import inspect import logging import pprint import signal @@ -21,6 +22,7 @@ from scrapy.extension import ExtensionManager from scrapy.interfaces import ISpiderLoader from scrapy.settings import overridden_settings, Settings from scrapy.signalmanager import SignalManager +from scrapy.utils.defer import deferred_from_coro from scrapy.utils.log import ( configure_logging, get_scrapy_root_handler, @@ -84,7 +86,10 @@ class Crawler: try: self.spider = self._create_spider(*args, **kwargs) self.engine = self._create_engine() - start_requests = iter(self.spider.start_requests()) + if inspect.iscoroutinefunction(self.spider.start_requests): + start_requests = yield deferred_from_coro(self.spider.start_requests()) + else: + start_requests = iter(self.spider.start_requests()) yield self.engine.open_spider(self.spider, start_requests) yield defer.maybeDeferred(self.engine.start) except Exception: diff --git a/tests/test_engine.py b/tests/test_engine.py index 5b7a4e676..4944423b9 100644 --- a/tests/test_engine.py +++ b/tests/test_engine.py @@ -15,6 +15,7 @@ import re import sys from urllib.parse import urlparse +from pytest import mark from twisted.internet import reactor, defer from twisted.web import server, static, util from twisted.trial import unittest @@ -83,6 +84,11 @@ class ItemZeroDivisionErrorSpider(TestSpider): } +class StartRequestsAsyncDefSpider(TestSpider): + async def start_requests(self): + return [Request(url, dont_filter=True) for url in self.start_urls] + + def start_test_site(debug=False): root_dir = os.path.join(tests_datadir, "test_site") r = static.File(root_dir) @@ -201,6 +207,13 @@ class EngineTest(unittest.TestCase): yield self.run.run() self._assert_items_error() + @mark.only_asyncio() + @defer.inlineCallbacks + def test_crawler_startrequests_asyncdef(self): + self.run = CrawlerRun(StartRequestsAsyncDefSpider) + yield self.run.run() + self._assert_visited_urls() + def _assert_visited_urls(self): must_be_visited = ["/", "/redirect", "/redirected", "/item1.html", "/item2.html", "/item999.html"]