diff --git a/CHANGELOG.md b/CHANGELOG.md index ff6f27ed8..5d87d0868 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,6 @@ # **Upcoming release** +- #802 Patch annotations on regular function arguments (@Strohutt) - ... # Release 1.15.0 diff --git a/rope/refactor/patchedast.py b/rope/refactor/patchedast.py index 2ea194d16..4112b01bd 100644 --- a/rope/refactor/patchedast.py +++ b/rope/refactor/patchedast.py @@ -632,7 +632,10 @@ def _NameConstant(self, node): self._handle(node, [str(node.value)]) def _arg(self, node): - self._handle(node, [node.arg]) + children = [node.arg] + if node.annotation is not None: + children.extend([":", node.annotation]) + self._handle(node, children) def _Pass(self, node): self._handle(node, ["pass"]) diff --git a/ropetest/refactor/patchedasttest.py b/ropetest/refactor/patchedasttest.py index 6998b24f7..b96778e3c 100644 --- a/ropetest/refactor/patchedasttest.py +++ b/ropetest/refactor/patchedasttest.py @@ -688,6 +688,36 @@ async def f(): ["async", " ", "def", " ", "f", "", "(", "", "arguments", "", ")", "", ":", "\n ", "Pass"], ) + def test_argument_annotation(self): + source = "def f(items: list):\n pass\n" + ast_frag = patchedast.get_patched_ast(source, True) + checker = _ResultChecker(self, ast_frag.body[0].args.args[0]) + start = source.index("items") + end = source.index(")") + checker.check_region("arg", start, end) + start = source.index("list") + checker.check_region("Name", start, start + len("list")) + checker.check_children("arg", ["items", "", ":", " ", "Name"]) + self.assertEqual(source, patchedast.write_ast(ast_frag)) + + def test_argument_annotation_with_default(self): + source = dedent("""\ + async def f( + items : dict[str, list[int]] = None, # keep the annotation + fallback=None, + ): + pass + """) + ast_frag = patchedast.get_patched_ast(source, True) + arguments = ast_frag.body[0].args + annotation = arguments.args[0].annotation + start = source.index("dict") + end = start + len("dict[str, list[int]]") + self.assertEqual((start, end), annotation.region) + self.assertEqual((source.index("items"), end), arguments.args[0].region) + self.assertEqual("None", source[slice(*arguments.defaults[0].region)]) + self.assertEqual(source, patchedast.write_ast(ast_frag)) + def test_function_node2(self): source = dedent('''\ def f(p1, **p2): diff --git a/ropetest/refactor/restructuretest.py b/ropetest/refactor/restructuretest.py index 9db6b77f7..f26af3fec 100644 --- a/ropetest/refactor/restructuretest.py +++ b/ropetest/refactor/restructuretest.py @@ -28,6 +28,16 @@ def test_replacing_simple_patterns(self): self.project.do(refactoring.get_changes()) self.assertEqual("a = int(1)\nb = 1\n", self.mod.read()) + def test_replacing_argument_annotations(self): + source = "def f(items: list[list[int]] = None):\n return items\n" + self.mod.write(source) + refactoring = restructure.Restructure(self.project, "list", "tuple") + self.project.do(refactoring.get_changes()) + self.assertEqual( + "def f(items: tuple[tuple[int]] = None):\n return items\n", + self.mod.read(), + ) + def test_replacing_patterns_with_normal_names(self): refactoring = restructure.Restructure( self.project, "${a} = 1", "${a} = int(1)", args={"a": "exact"} diff --git a/ropetest/refactor/similarfindertest.py b/ropetest/refactor/similarfindertest.py index a4eed44d3..2ae6584f5 100644 --- a/ropetest/refactor/similarfindertest.py +++ b/ropetest/refactor/similarfindertest.py @@ -24,6 +24,21 @@ def test_trivial_case(self): finder = self._create_finder("") self.assertEqual([], list(finder.get_match_regions("10"))) + def test_matching_argument_annotation(self): + source = "def foo(l: list):\n pass\n" + finder = similarfinder.RawSimilarFinder(source) + regions = [match.get_region() for match in finder.get_matches("${a}")] + start = source.index("list") + self.assertEqual([(start, start + len("list"))], regions) + + def test_matching_nested_argument_annotation(self): + source = "def foo(items: list[list[int]] = None):\n pass\n" + finder = self._create_finder(source) + regions = list(finder.get_match_regions("list")) + first = source.index("list") + second = source.index("list", first + 1) + self.assertEqual([(first, first + 4), (second, second + 4)], regions) + def test_constant_integer(self): source = "a = 10\n" finder = self._create_finder(source)