From 21d3187943ec08ceb5727bfe64bdc18304c2a2cc Mon Sep 17 00:00:00 2001 From: ethanstoner Date: Wed, 16 Sep 2026 19:07:26 -0700 Subject: [PATCH 1/2] Define functions extracted from a class body before the class Extract Method on code in a class body inserted the new function after the class, but class bodies run at definition time, so the class raised NameError. Insert the function before the (outermost) class instead, and collect class-body variables so they are passed as arguments. Fixes #825 --- CHANGELOG.md | 1 + rope/refactor/extract.py | 41 ++++++++- ropetest/refactor/extracttest.py | 147 +++++++++++++++++++++++++++++++ 3 files changed, 186 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1650aa77a..d34d86877 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 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 1b1659fad..97c16bae6 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 4b34d58d3..60b0a2044 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(): From 2a179e7cd0f930cf021f64d450174903ad6f2908 Mon Sep 17 00:00:00 2001 From: ethanstoner Date: Wed, 16 Sep 2026 19:07:44 -0700 Subject: [PATCH 2/2] Add PR number to CHANGELOG --- CHANGELOG.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d34d86877..55d38f88f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,7 +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 Define functions extracted from a class body before the class (@ethanstoner) +- #825, #878 Define functions extracted from a class body before the class (@ethanstoner) # Release 1.14.0