diff --git a/scrapy/utils/defer.py b/scrapy/utils/defer.py index 6db9cc117..2d02c0621 100644 --- a/scrapy/utils/defer.py +++ b/scrapy/utils/defer.py @@ -124,6 +124,20 @@ def iter_errback(iterable, errback, *a, **kw): errback(failure.Failure(), *a, **kw) +async def aiter_errback(aiterable, errback, *a, **kw): + """Wraps an async iterable calling an errback if an error is caught while + iterating it. Similar to scrapy.utils.defer.iter_errback() + """ + it = aiterable.__aiter__() + while True: + try: + yield await it.__anext__() + except StopAsyncIteration: + break + except Exception: + errback(failure.Failure(), *a, **kw) + + def deferred_from_coro(o): """Converts a coroutine into a Deferred, or returns the object as is if it isn't a coroutine""" if isinstance(o, defer.Deferred): diff --git a/tests/test_utils_defer.py b/tests/test_utils_defer.py index e60242a3b..06d91c574 100644 --- a/tests/test_utils_defer.py +++ b/tests/test_utils_defer.py @@ -2,12 +2,15 @@ from twisted.trial import unittest from twisted.internet import reactor, defer from twisted.python.failure import Failure +from scrapy.utils.asyncgen import collect_asyncgen from scrapy.utils.defer import ( iter_errback, + aiter_errback, mustbe_deferred, process_chain, process_chain_both, process_parallel, + deferred_f_from_coro_f, ) @@ -117,3 +120,31 @@ class IterErrbackTest(unittest.TestCase): self.assertEqual(out, [0, 1, 2, 3, 4]) self.assertEqual(len(errors), 1) self.assertIsInstance(errors[0].value, ZeroDivisionError) + + +class AiterErrbackTest(unittest.TestCase): + + @deferred_f_from_coro_f + async def test_aiter_errback_good(self): + async def itergood(): + for x in range(10): + yield x + + errors = [] + out = await collect_asyncgen(aiter_errback(itergood(), errors.append)) + self.assertEqual(out, list(range(10))) + self.assertFalse(errors) + + @deferred_f_from_coro_f + async def test_iter_errback_bad(self): + async def iterbad(): + for x in range(10): + if x == 5: + 1 / 0 + yield x + + errors = [] + out = await collect_asyncgen(aiter_errback(iterbad(), errors.append)) + self.assertEqual(out, [0, 1, 2, 3, 4]) + self.assertEqual(len(errors), 1) + self.assertIsInstance(errors[0].value, ZeroDivisionError)