mirror of https://github.com/scrapy/scrapy.git
Complete test coverage for exception scenarios of SpiderMiddleware.process_seeds
This commit is contained in:
parent
b587b00439
commit
0a137e40b7
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue