diff --git a/CHANGELOG.md b/CHANGELOG.md index 1650aa77..55d38f88 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,7 @@ - #623, #819, #863 Support MatchOr, MatchSequence, MatchStar (@jheld, @lieryan) - #870 Add default implementation for is_dir() (@lieryan) - #872 Fix unicode handling in patchedast (@lieryan) +- #825, #878 Define functions extracted from a class body before the class (@ethanstoner) # Release 1.14.0 diff --git a/rope/refactor/extract.py b/rope/refactor/extract.py index 1b1659fa..97c16bae 100644 --- a/rope/refactor/extract.py +++ b/rope/refactor/extract.py @@ -206,7 +206,27 @@ def global_(self): @property def method(self): - return self.scope.parent is not None and self.scope.parent.get_kind() == "Class" + return ( + self.scope.get_kind() == "Function" + and self.scope.parent is not None + and self.scope.parent.get_kind() == "Class" + ) + + @property + def class_to_define_before(self): + """The class the extracted function must be defined before, if any + + Class bodies run when the class is defined, so a function extracted + from a class body has to be defined before the (outermost) class. + """ + if self.scope.get_kind() != "Class": + return None + scope = self.scope + while scope.parent.get_kind() == "Class": + scope = scope.parent + if self.make_global and scope.parent.get_kind() != "Module": + return None + return scope @property def indents(self): @@ -400,6 +420,8 @@ def __init__(self, info, matched_lines): def find_lineno(self): if self.info.variable and not self.info.make_global: return self._get_before_line() + if self.info.class_to_define_before is not None: + return self._get_before_class() if self.info.global_: toplevel = self._find_toplevel(self.info.scope) ast = self.info.pymodule.get_ast() @@ -420,6 +442,10 @@ def _find_toplevel(self, scope): def find_indents(self): if self.info.variable and not self.info.make_global: return sourceutils.get_indents(self.info.lines, self._get_before_line()) + if self.info.class_to_define_before is not None: + return sourceutils.get_indents( + self.info.lines, self.info.class_to_define_before.get_start() + ) else: if self.info.global_ or self.info.make_global: return 0 @@ -432,6 +458,10 @@ def _get_before_line(self): def _get_after_scope(self): return self.info.scope.get_end() + 1 + def _get_before_class(self): + node = self.info.class_to_define_before.pyobject.get_ast() + return min([node.lineno] + [d.lineno for d in node.decorator_list]) + class _ExceptionalConditionChecker: def __call__(self, info): @@ -554,7 +584,7 @@ def _extracting_from_classmethod(self): return self.info.method and _get_function_kind(self.info.scope) == "classmethod" def get_definition(self): - if self.info.global_: + if self.info.global_ or self.info.class_to_define_before is not None: return "\n%s\n" % self._get_function_definition() else: return "\n%s" % self._get_function_definition() @@ -900,7 +930,12 @@ def _AugAssign(self, node): self.visit(node.target) def _ClassDef(self, node): - self._written_variable(node.name, node.lineno) + if not self.is_global and self.host_function: + self.host_function = False + for child in node.body: + self.visit(child) + else: + self._written_variable(node.name, node.lineno) def _ListComp(self, node): self._comp_exp(node) diff --git a/ropetest/refactor/extracttest.py b/ropetest/refactor/extracttest.py index 4b34d58d..60b0a204 100644 --- a/ropetest/refactor/extracttest.py +++ b/ropetest/refactor/extracttest.py @@ -1758,6 +1758,153 @@ def new_func(): """) self.assertEqual(expected, refactored) + def test_extract_method_in_class_body(self): + code = dedent("""\ + class TSV: + delimiter = "\\t" + + print(TSV.delimiter) + """) + start = code.index('"\\t"') + end = start + len('"\\t"') + refactored = self.do_extract_method(code, start, end, "extracted") + expected = dedent("""\ + + def extracted(): + return "\\t" + + class TSV: + delimiter = extracted() + + print(TSV.delimiter) + """) + self.assertEqual(expected, refactored) + + def test_extract_method_in_class_body_reading_class_variables(self): + code = dedent("""\ + class A: + a = 1 + b = a + 2 + """) + start = code.index("a + 2") + end = start + len("a + 2") + refactored = self.do_extract_method(code, start, end, "extracted") + expected = dedent("""\ + + def extracted(a): + return a + 2 + + class A: + a = 1 + b = extracted(a) + """) + self.assertEqual(expected, refactored) + + def test_extract_method_in_decorated_class_body(self): + code = dedent("""\ + @decorator + class A: + a = 1 + 2 + """) + start = code.index("1 + 2") + end = start + len("1 + 2") + refactored = self.do_extract_method(code, start, end, "extracted") + expected = dedent("""\ + + def extracted(): + return 1 + 2 + + @decorator + class A: + a = extracted() + """) + self.assertEqual(expected, refactored) + + def test_extract_method_in_class_body_inside_function(self): + code = dedent("""\ + def f(): + class A: + a = 1 + 2 + return A + """) + start = code.index("1 + 2") + end = start + len("1 + 2") + refactored = self.do_extract_method(code, start, end, "extracted") + expected = dedent("""\ + def f(): + + def extracted(): + return 1 + 2 + + class A: + a = extracted() + return A + """) + self.assertEqual(expected, refactored) + + def test_extract_method_in_nested_class_body(self): + code = dedent("""\ + class A: + class B: + a = 1 + 2 + """) + start = code.index("1 + 2") + end = start + len("1 + 2") + refactored = self.do_extract_method(code, start, end, "extracted") + expected = dedent("""\ + + def extracted(): + return 1 + 2 + + class A: + class B: + a = extracted() + """) + self.assertEqual(expected, refactored) + + def test_global_extract_method_in_class_body(self): + code = dedent("""\ + def f(): + class A: + a = 1 + 2 + return A + """) + start = code.index("1 + 2") + end = start + len("1 + 2") + refactored = self.do_extract_method( + code, start, end, "extracted", global_=True + ) + expected = dedent("""\ + def f(): + class A: + a = extracted() + return A + + def extracted(): + return 1 + 2 + """) + self.assertEqual(expected, refactored) + + def test_global_extract_method_in_module_level_class_body(self): + code = dedent("""\ + class A: + a = 1 + 2 + """) + start = code.index("1 + 2") + end = start + len("1 + 2") + refactored = self.do_extract_method( + code, start, end, "extracted", global_=True + ) + expected = dedent("""\ + + def extracted(): + return 1 + 2 + + class A: + a = extracted() + """) + self.assertEqual(expected, refactored) + def test_where_to_search_when_extracting_global_names(self): code = dedent("""\ def a():