From d68aab992e435aa1c88061e83d03e57e7d31533d Mon Sep 17 00:00:00 2001 From: Grisha Temchenko Date: Thu, 20 Aug 2020 09:22:07 -0400 Subject: [PATCH] Smarter generator check for combined return/yield statements (#4721) --- scrapy/utils/misc.py | 19 ++++++++++++++++++- ...t_return_with_argument_inside_generator.py | 19 +++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/scrapy/utils/misc.py b/scrapy/utils/misc.py index d6966be8e..bd400bd30 100644 --- a/scrapy/utils/misc.py +++ b/scrapy/utils/misc.py @@ -5,6 +5,7 @@ import os import re import hashlib import warnings +from collections import deque from contextlib import contextmanager from importlib import import_module from pkgutil import iter_modules @@ -184,6 +185,22 @@ def set_environ(**kwargs): os.environ[k] = v +def walk_callable(node): + """Similar to ``ast.walk``, but walks only function body and skips nested + functions defined within the node. + """ + todo = deque([node]) + walked_func_def = False + while todo: + node = todo.popleft() + if isinstance(node, ast.FunctionDef): + if walked_func_def: + continue + walked_func_def = True + todo.extend(ast.iter_child_nodes(node)) + yield node + + _generator_callbacks_cache = LocalWeakReferencedCache(limit=128) @@ -201,7 +218,7 @@ def is_generator_with_return_value(callable): if inspect.isgeneratorfunction(callable): tree = ast.parse(dedent(inspect.getsource(callable))) - for node in ast.walk(tree): + for node in walk_callable(tree): if isinstance(node, ast.Return) and not returns_none(node): _generator_callbacks_cache[callable] = True return _generator_callbacks_cache[callable] diff --git a/tests/test_utils_misc/test_return_with_argument_inside_generator.py b/tests/test_utils_misc/test_return_with_argument_inside_generator.py index bdbec1beb..2be38620c 100644 --- a/tests/test_utils_misc/test_return_with_argument_inside_generator.py +++ b/tests/test_utils_misc/test_return_with_argument_inside_generator.py @@ -29,9 +29,28 @@ class UtilsMiscPy3TestCase(unittest.TestCase): yield 1 yield from g() + def m(): + yield 1 + + def helper(): + return 0 + + yield helper() + + def n(): + yield 1 + + def helper(): + return 0 + + yield helper() + return 2 + assert is_generator_with_return_value(f) assert is_generator_with_return_value(g) assert not is_generator_with_return_value(h) assert not is_generator_with_return_value(i) assert not is_generator_with_return_value(j) assert not is_generator_with_return_value(k) # not recursive + assert not is_generator_with_return_value(m) + assert is_generator_with_return_value(n)