Complete test coverage for exception scenarios of SpiderMiddleware.process_seeds

This commit is contained in:
Adrián Chaves 2025-03-15 00:55:31 +01:00
parent b587b00439
commit 0a137e40b7
2 changed files with 87 additions and 7 deletions

View File

@ -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):

View File

@ -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