Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
41 changes: 38 additions & 3 deletions rope/refactor/extract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand All @@ -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):
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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)
Expand Down
147 changes: 147 additions & 0 deletions ropetest/refactor/extracttest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down