From c961438d5d9998344460d930ea502fae40553043 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Chaves?= Date: Wed, 16 Mar 2022 18:45:56 +0100 Subject: [PATCH] tests: cover scenarios of bad results from process_spider_output --- scrapy/core/spidermw.py | 19 ++++++++---- tests/test_spidermiddleware.py | 54 +++++++++++++++++++++++++++------- 2 files changed, 58 insertions(+), 15 deletions(-) diff --git a/scrapy/core/spidermw.py b/scrapy/core/spidermw.py index 6075670b0..1aa02f29f 100644 --- a/scrapy/core/spidermw.py +++ b/scrapy/core/spidermw.py @@ -4,7 +4,7 @@ Spider Middleware manager See documentation in docs/topics/spider-middleware.rst """ import logging -from inspect import isasyncgenfunction +from inspect import isasyncgenfunction, iscoroutine from itertools import islice from typing import Any, AsyncGenerator, AsyncIterable, Callable, Generator, Iterable, Tuple, Union, cast @@ -61,7 +61,7 @@ class SpiderMiddlewareManager(MiddlewareManager): try: result = method(response=response, spider=spider) if result is not None: - msg = (f"Middleware {method.__qualname__} must return None " + msg = (f"{method.__qualname__} must return None " f"or raise an exception, got {type(result)}") raise _InvalidOutput(msg) except _InvalidOutput: @@ -129,7 +129,7 @@ class SpiderMiddlewareManager(MiddlewareManager): elif result is None: continue else: - msg = (f"Middleware {method.__qualname__} must return None " + msg = (f"{method.__qualname__} must return None " f"or an iterable, got {type(result)}") raise _InvalidOutput(msg) return _failure @@ -197,8 +197,17 @@ class SpiderMiddlewareManager(MiddlewareManager): if _isiterable(result): result = self._evaluate_iterable(response, spider, result, method_index + 1, recovered) else: - msg = (f"Middleware {method.__qualname__} must return an " - f"iterable, got {type(result)}") + if iscoroutine(result): + result.close() # Silence warning about not awaiting + msg = ( + f"{method.__qualname__} must be an asynchronous " + f"generator (i.e. use yield)" + ) + else: + msg = ( + f"{method.__qualname__} must return an iterable, got " + f"{type(result)}" + ) raise _InvalidOutput(msg) last_result_is_async = isinstance(result, AsyncIterable) diff --git a/tests/test_spidermiddleware.py b/tests/test_spidermiddleware.py index f9f2b6642..ed0912b82 100644 --- a/tests/test_spidermiddleware.py +++ b/tests/test_spidermiddleware.py @@ -123,6 +123,11 @@ class BaseAsyncSpiderMiddlewareTestCase(SpiderMiddlewareTestCase): start_index = 10 return {i: c for c, i in enumerate(mw_classes, start=start_index)} + def _scrape_func(self, *args, **kwargs): + yield {'foo': 1} + yield {'foo': 2} + yield {'foo': 3} + @defer.inlineCallbacks def _get_middleware_result(self, *mw_classes, start_index: Optional[int] = None): setting = self._construct_mw_setting(*mw_classes, start_index=start_index) @@ -201,11 +206,6 @@ class ProcessSpiderOutputSimple(BaseAsyncSpiderMiddlewareTestCase): MW_ASYNCGEN = ProcessSpiderOutputAsyncGenMiddleware MW_UNIVERSAL = ProcessSpiderOutputUniversalMiddleware - def _scrape_func(self, *args, **kwargs): - yield {'foo': 1} - yield {'foo': 2} - yield {'foo': 3} - def test_simple(self): """ Simple mw """ return self._test_simple_base(self.MW_SIMPLE) @@ -285,6 +285,45 @@ class ProcessSpiderOutputAsyncGen(ProcessSpiderOutputSimple): downgrade=True) +class ProcessSpiderOutputNonIterableMiddleware: + def process_spider_output(self, response, result, spider): + return + + +class ProcessSpiderOutputCoroutineMiddleware: + async def process_spider_output(self, response, result, spider): + results = [] + for r in result: + results.append(r) + return results + + +class ProcessSpiderOutputInvalidResult(BaseAsyncSpiderMiddlewareTestCase): + + @defer.inlineCallbacks + def test_non_iterable(self): + with self.assertRaisesRegex( + _InvalidOutput, + ( + "\.process_spider_output must return an iterable, got " + ), + ): + yield self._get_middleware_result( + ProcessSpiderOutputNonIterableMiddleware, + ) + + @defer.inlineCallbacks + def test_coroutine(self): + with self.assertRaisesRegex( + _InvalidOutput, + "\.process_spider_output must be an asynchronous generator", + ): + yield self._get_middleware_result( + ProcessSpiderOutputCoroutineMiddleware, + ) + + class ProcessStartRequestsSimpleMiddleware: def process_start_requests(self, start_requests, spider): for r in start_requests: @@ -387,11 +426,6 @@ class BuiltinMiddlewareSimpleTest(BaseAsyncSpiderMiddlewareTestCase): MW_ASYNCGEN = ProcessSpiderOutputAsyncGenMiddleware MW_UNIVERSAL = ProcessSpiderOutputUniversalMiddleware - def _scrape_func(self, *args, **kwargs): - yield {'foo': 1} - yield {'foo': 2} - yield {'foo': 3} - @defer.inlineCallbacks def _get_middleware_result(self, *mw_classes, start_index: Optional[int] = None): setting = self._construct_mw_setting(*mw_classes, start_index=start_index)