From 0a137e40b7091b8df4118f1ec6c718a51df6ee4b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Chaves?= Date: Sat, 15 Mar 2025 00:55:31 +0100 Subject: [PATCH] Complete test coverage for exception scenarios of SpiderMiddleware.process_seeds --- scrapy/core/spidermw.py | 22 ++++-- tests/test_spidermiddleware_process_seeds.py | 72 ++++++++++++++++++++ 2 files changed, 87 insertions(+), 7 deletions(-) diff --git a/scrapy/core/spidermw.py b/scrapy/core/spidermw.py index c949c3784..f94b8025f 100644 --- a/scrapy/core/spidermw.py +++ b/scrapy/core/spidermw.py @@ -372,7 +372,11 @@ class SpiderMiddlewareManager(MiddlewareManager): def process_seeds( self, spider: Spider ) -> Generator[Deferred[Any], Any, AsyncIterator[Any] | None]: - self._check_deprecated_start_requests_use(spider) + try: + self._check_deprecated_start_requests_use(spider) + except ValueError as exception: + logger.error(exception) + return None seeds: AsyncIterator[Any] if self._use_start_requests: sync_seeds = iter(spider.start_requests()) @@ -381,15 +385,19 @@ class SpiderMiddlewareManager(MiddlewareManager): ) seeds = as_async_generator(sync_seeds) else: - if not isasyncgenfunction(spider.yield_seeds): + error_found = False + for fn in (spider.yield_seeds, *self.methods["process_seeds"]): + if isasyncgenfunction(fn): + continue logger.error( - f"{global_object_name(spider.yield_seeds)} must be an " - f"async generator function, i.e. an async def function " - f"with yield statements." + f"{global_object_name(fn)} must be an async generator " + f"function, i.e. an async def function with yield " + f"statements." ) + error_found = True + if error_found: return None - seeds = spider.yield_seeds() - seeds = yield self._process_chain("process_seeds", seeds) + seeds = yield self._process_chain("process_seeds", spider.yield_seeds()) return seeds def _check_deprecated_start_requests_use(self, spider: Spider): diff --git a/tests/test_spidermiddleware_process_seeds.py b/tests/test_spidermiddleware_process_seeds.py index f7bac17de..fccdca0f1 100644 --- a/tests/test_spidermiddleware_process_seeds.py +++ b/tests/test_spidermiddleware_process_seeds.py @@ -102,6 +102,8 @@ class DeprecatedWrapSpiderMiddleware: class MainTestCase(TestCase): + # Helper methods + async def _test(self, spider_middlewares, spider_cls, expected_items): actual_items = [] @@ -117,6 +119,17 @@ class MainTestCase(TestCase): assert crawler.stats.get_value("finish_reason") == "finished" assert actual_items == expected_items, f"{actual_items=} != {expected_items=}" + async def _test_process_seeds(self, _process_seeds, expected_items=None): + class TestSpiderMiddleware: + process_seeds = _process_seeds + + class TestSpider(Spider): + name = "test" + + await self._test([TestSpiderMiddleware], TestSpider, expected_items) + + # Deprecation and universal + async def _test_wrap(self, spider_middleware, spider_cls, expected_items=None): expected_items = ( [ITEM_A, ITEM_B, ITEM_C] if expected_items is None else expected_items @@ -190,6 +203,8 @@ class MainTestCase(TestCase): ): await self._test_wrap(DeprecatedWrapSpiderMiddleware, DeprecatedWrapSpider) + # Sleep tests + async def _test_sleep(self, spider_middlewares): class TestSpider(Spider): name = "test" @@ -220,3 +235,60 @@ class MainTestCase(TestCase): await self._test_sleep( [NoOpSpiderMiddleware, TwistedSleepSpiderMiddleware, NoOpSpiderMiddleware] ) + + # Bad definitions. + + @deferred_f_from_coro_f + async def test_async_function(self): + async def process_seeds(mw, seeds): + return + + with LogCapture() as log: + await self._test_process_seeds(process_seeds, []) + + assert ".process_seeds must be an async generator function" in str(log), log + + @deferred_f_from_coro_f + async def test_sync_function(self): + def process_seeds(mw, spider): + return [] + + with LogCapture() as log: + await self._test_process_seeds(process_seeds, []) + + assert ".process_seeds must be an async generator function" in str(log) + + @deferred_f_from_coro_f + async def test_sync_generator(self): + def process_seeds(mw, spider): + return + yield + + with LogCapture() as log: + await self._test_process_seeds(process_seeds, []) + + assert ".process_seeds must be an async generator function" in str(log) + + # Exceptions during iteration. + + @deferred_f_from_coro_f + async def test_exception_before_yield(self): + async def process_seeds(mw, seeds): + raise RuntimeError + yield + + with LogCapture() as log: + await self._test_process_seeds(process_seeds, []) + + assert "in process_seeds\n raise RuntimeError" in str(log), log + + @deferred_f_from_coro_f + async def test_exception_after_yield(self): + async def process_seeds(mw, spider): + yield ITEM_A + raise RuntimeError + + with LogCapture() as log: + await self._test_process_seeds(process_seeds, [ITEM_A]) + + assert "in process_seeds\n raise RuntimeError" in str(log), log