From 777bfc96182b24d3dba9c49f5aa551f783d3d4f4 Mon Sep 17 00:00:00 2001 From: Strohutt <84200776+Strohutt@users.noreply.github.com> Date: Sun, 27 Sep 2026 16:36:46 +0200 Subject: [PATCH] Fix source regions for regular argument annotations Visit each argument's annotation so structural searches and replacements can use its source region. Cover the reported wildcard crash, nested annotations, defaults, and source round trips. Based on the diagnosis and proposed fix in issue #802. --- CHANGELOG.md | 1 + rope/refactor/patchedast.py | 5 ++++- ropetest/refactor/patchedasttest.py | 30 ++++++++++++++++++++++++++ ropetest/refactor/restructuretest.py | 10 +++++++++ ropetest/refactor/similarfindertest.py | 15 +++++++++++++ 5 files changed, 60 insertions(+), 1 deletion(-) 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)