tests: cover scenarios of bad results from process_spider_output

This commit is contained in:
Adrián Chaves 2022-03-16 18:45:56 +01:00
parent fd08bb6cd9
commit c961438d5d
2 changed files with 58 additions and 15 deletions

View File

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

View File

@ -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 <class "
"'NoneType'>"
),
):
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)