diff --git a/CHANGELOG.md b/CHANGELOG.md index b3e7311e..7cbb8d46 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,6 @@ # **Upcoming release** +- Refuse variable inlining that moves an initializer into a deferred annotation or a lazy function/class type parameter. - #895 patchedast Starred and keyword now consumes their syntactically expected and ** (@lieryan) - #896 patchedast cleanup and refactoring (@lieryan) diff --git a/rope/refactor/inline.py b/rope/refactor/inline.py index 768d30b7..a8ec58e2 100644 --- a/rope/refactor/inline.py +++ b/rope/refactor/inline.py @@ -17,10 +17,12 @@ # but it should be 200. import re +import sys from typing import List import rope.base.builtins # Use fully qualified names for clarity. from rope.base import ( + ast, codeanalyze, evaluate, exceptions, @@ -267,6 +269,14 @@ def get_changes( resources = [self.original] if remove and self.original != self.resource: resources.append(self.resource) + resources = list(resources) + check_resources = resources + if remove and not rename._is_local(self.pyname): + check_resources = list( + dict.fromkeys(resources + list(self.project.get_python_files())) + ) + for resource in check_resources: + self._check_deferred_annotation_references(resource, remove, only_current) changes = ChangeSet("Inline variable <%s>" % self.name) jobset = task_handle.create_jobset("Calculating changes", len(resources)) @@ -283,6 +293,203 @@ def get_changes( jobset.finished_job() return changes + def _check_deferred_annotation_references(self, resource, remove, only_current): + pymodule = self.project.get_pymodule(resource) + tree = pymodule.get_ast() + future_annotations = any( + isinstance(node, ast.ImportFrom) + and node.module == "__future__" + and any(alias.name == "annotations" for alias in node.names) + for node in tree.body + ) + if not future_annotations and sys.version_info < (3, 12): + return + annotations = [] + type_parameters = [] + + def collect(node, scope=None): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + parameters = getattr(node, "type_params", []) + if parameters: + namespace = ( + pymodule.get_scope() + .get_inner_scope_for_line(node.lineno) + .parent + ) + names = {parameter.name for parameter in parameters} + for parameter in parameters: + for attr in ("bound", "default_value"): + expression = getattr(parameter, attr, None) + if expression is not None: + type_parameters.append((expression, namespace, names)) + if isinstance(node, ast.arg): + annotations.append(node.annotation) + elif isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + annotations.append(node.returns) + scope = "function" + elif isinstance(node, ast.Lambda): + scope = "function" + elif isinstance(node, ast.ClassDef): + scope = "class" + elif ( + isinstance(node, ast.AnnAssign) and node.simple and scope != "function" + ): + annotations.append(node.annotation) + for child in ast.iter_child_nodes(node): + collect(child, scope) + + collect(tree) + if not future_annotations and sys.version_info < (3, 14): + annotations = [] + lines = codeanalyze.ASTLinesAdapter(pymodule.source_code) + regions = [lines[node] for node in annotations if node is not None] + parameter_regions = [lines[item[0]] for item in type_parameters] + if not regions and not parameter_regions: + return + + def parameter_reference(node, namespace, names): + inner_namespace = namespace + inner_names = set(names) + while inner_namespace.get_kind() == "Class": + inner_names.update( + parameter.name + for parameter in getattr( + inner_namespace.pyobject.get_ast(), "type_params", [] + ) + ) + inner_namespace = inner_namespace.parent + if isinstance(node, ast.Lambda): + arguments = node.args + defaults = arguments.defaults + [ + default for default in arguments.kw_defaults if default is not None + ] + bound_names = inner_names | { + argument.arg + for argument in ( + arguments.posonlyargs + + arguments.args + + arguments.kwonlyargs + + [arguments.vararg, arguments.kwarg] + ) + if argument is not None + } + return any( + parameter_reference(default, namespace, names) + for default in defaults + ) or parameter_reference(node.body, inner_namespace, bound_names) + if isinstance( + node, (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp) + ): + bound_names = set(names) + for index, generator in enumerate(node.generators): + iterable_namespace = namespace if index == 0 else inner_namespace + if parameter_reference( + generator.iter, iterable_namespace, bound_names + ): + return True + bound_names.update(inner_names) + bound_names.update( + name.id + for name in ast.walk(generator.target) + if isinstance(name, ast.Name) + and isinstance(name.ctx, ast.Store) + ) + if parameter_reference( + generator.target, inner_namespace, bound_names + ): + return True + if any( + parameter_reference(condition, inner_namespace, bound_names) + for condition in generator.ifs + ): + return True + values = ( + [node.key, node.value] + if isinstance(node, ast.DictComp) + else [node.elt] + ) + return any( + parameter_reference(value, inner_namespace, bound_names) + for value in values + ) + if isinstance(node, (ast.Name, ast.Attribute)): + primary = node + while isinstance(primary, ast.Attribute): + primary = primary.value + if not (isinstance(primary, ast.Name) and primary.id in names): + if isinstance(primary, ast.Name): + outer_scope = namespace + while outer_scope is not None: + if ( + outer_scope.get_kind() != "Class" + or outer_scope is namespace + ) and primary.id in outer_scope.get_names(): + break + if any( + parameter.name == primary.id + for parameter in getattr( + outer_scope.pyobject.get_ast(), "type_params", [] + ) + ): + return False + outer_scope = outer_scope.parent + if occurrences.same_pyname( + self.pyname, evaluate.eval_node(namespace, node) + ): + return True + return any( + parameter_reference(child, namespace, names) + for child in ast.iter_child_nodes(node) + ) + + if remove and any( + parameter_reference(expression, namespace, names) + for expression, namespace, names in type_parameters + ): + # Type parameter expressions use the enclosing annotation scope, + # not the function's parameters/locals or the new class's body. + raise exceptions.RefactoringError( + "Cannot inline a variable referenced in a deferred annotation." + ) + if ( + future_annotations + and remove + and any( + isinstance(node, (ast.Name, ast.Attribute)) + and occurrences.same_pyname( + self.pyname, evaluate.eval_node(pymodule.get_scope(), node) + ) + for annotation in annotations + if annotation is not None + for node in ast.walk(annotation) + ) + ): + # Stringified annotations may resolve module globals even when + # static name lookup finds a class member or function parameter. + raise exceptions.RefactoringError( + "Cannot inline a variable referenced in a deferred annotation." + ) + finder = occurrences.create_finder( + self.project, self.name, self.pyname, imports=False + ) + for occurrence in finder.find_occurrences(pymodule=pymodule): + current = True + if only_current: + start, end = occurrence.get_primary_range() + current = resource == self.original and start <= self.offset <= end + if not remove and not current: + continue + start, _ = occurrence.get_word_range() + if any(begin <= start < end for begin, end in regions) or ( + current + and any(begin <= start < end for begin, end in parameter_regions) + ): + # Substituting the initializer can change when an annotation + # evaluates it and which bindings it observes. + raise exceptions.RefactoringError( + "Cannot inline a variable referenced in a deferred annotation." + ) + def _change_main_module(self, remove, only_current, docs): region = None if only_current and self.original == self.resource: diff --git a/ropetest/refactor/inline_annotations_test.py b/ropetest/refactor/inline_annotations_test.py new file mode 100644 index 00000000..3e4953c0 --- /dev/null +++ b/ropetest/refactor/inline_annotations_test.py @@ -0,0 +1,424 @@ +import subprocess +import sys +from textwrap import dedent + +import pytest + +from rope.base.exceptions import RefactoringError +from rope.base.project import Project +from rope.refactor.inline import create_inline + + +@pytest.fixture +def project(tmp_path): + project = Project( + str(tmp_path), save_objectdb=False, save_history=False, automatic_soa=False + ) + yield project + project.close() + + +def module(project, name, source): + resource = project.root.create_file(name + ".py") + resource.write(dedent(source)) + return resource + + +def execute(project, resource): + result = subprocess.run( + [sys.executable, str(resource.real_path)], + cwd=project.address, + capture_output=True, + text=True, + timeout=10, + ) + assert result.returncode == 0, result.stderr + return result.stdout + + +@pytest.fixture(params=[True, False], ids=["future", "default"]) +def deferred_prefix(request): + if not request.param and sys.version_info < (3, 14): + pytest.skip("default annotations are deferred on Python 3.14+") + return "from __future__ import annotations\n" if request.param else "" + + +@pytest.mark.parametrize("remove", [True, False]) +def test_refuses_changing_annotation_name_binding(project, deferred_prefix, remove): + source = module(project, "source", deferred_prefix + dedent('''\ + original = int + target = original + def func(x: target): + pass + original = str + import typing + print(typing.get_type_hints(func)["x"] is int) + ''')) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes(remove=remove) + assert source.read() == before + assert execute(project, source) == output + + +def test_refuses_moving_side_effect_into_annotation(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + events = [] + def make(): + events.append("called") + return int + target = make() + def func() -> target: + pass + print(events) + import typing + print(typing.get_type_hints(func)["return"] is int) + print(events) + ''')) + before = source.read() + output = execute(project, source) + assert output == "['called']\nTrue\n['called']\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize("definition,owner,key", [ + (dedent("""\ + async def func(x: target): + pass + """), "func", "x"), + (dedent("""\ + def func(x: target, /): + pass + """), "func", "x"), + (dedent("""\ + def func(*, x: target): + pass + """), "func", "x"), + (dedent("""\ + def func(*x: target): + pass + """), "func", "x"), + (dedent("""\ + def func(**x: target): + pass + """), "func", "x"), + (dedent("""\ + def func() -> target: + pass + """), "func", "return"), + (dedent("""\ + class Owner: + x: target + """), "Owner", "x"), + (dedent("""\ + x: target + import sys + """), "sys.modules[__name__]", "x"), + (dedent("""\ + def 函数(x: ( + list[target] + )): + pass + """), "函数", "x"), +]) +def test_refuses_annotation_kinds(project, deferred_prefix, definition, owner, key): + comparison = "list[int]" if "list[target]" in definition else "int" + code = deferred_prefix + "original = int\ntarget = original\n" + definition + code += "original = str\nimport typing\n" + code += f'print(typing.get_type_hints({owner})["{key}"] == {comparison})\n' + source = module(project, "source", code) + assert execute(project, source) == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, code.index("target") + 1).get_changes() + assert source.read() == code + assert execute(project, source) == "True\n" + + +@pytest.mark.parametrize("qualified", [True, False]) +def test_cross_module_annotation_is_refused(project, deferred_prefix, qualified): + source = module(project, "source", "original = int\ntarget = original\n") + imported = "source.target" if qualified else "target" + import_line = "import source\n" if qualified else "from source import target\n" + code = deferred_prefix + import_line + code += f"def func(x: {imported}):\n pass\n" + code += 'import typing\nprint(typing.get_type_hints(func)["x"] is int)\n' + entry = module(project, "entry", code) + before = {source: source.read(), entry: entry.read()} + assert execute(project, entry) == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, entry, code.index("target", code.index("def func")) + 1).get_changes(only_current=True, remove=False) + assert {r: r.read() for r in before} == before + assert execute(project, entry) == "True\n" + + +@pytest.mark.parametrize("remove", [False, True]) +def test_only_current_eager_reference(project, deferred_prefix, remove): + source = module(project, "source", deferred_prefix + dedent('''\ + target = int + def func(x: target): + pass + result = target + import typing + print(result is typing.get_type_hints(func)["x"]) + ''')) + before = source.read() + output = execute(project, source) + inline = create_inline(project, source, before.rindex("target") + 1) + if remove: + with pytest.raises(RefactoringError, match="deferred annotation"): + inline.get_changes(only_current=True) + assert source.read() == before + else: + project.do(inline.get_changes(only_current=True, remove=False)) + assert "result = int" in source.read() + assert "target = int" in source.read() + assert execute(project, source) == output == "True\n" + + +def test_eager_annotations_allow_inline(project): + if sys.version_info >= (3, 14): + pytest.skip("default annotations are deferred on Python 3.14+") + source = module(project, "source", '''\ + original = int + target = original + def func(x: target): + pass + original = str + import typing + print(typing.get_type_hints(func)["x"] is int) + ''') + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "target" not in source.read() + assert execute(project, source) == before == "True\n" + + +def test_function_default_is_eager(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + original = int + target = original + def func(x: str = target): + return x + original = str + print(func() is int) + ''')) + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "x: str = original" in source.read() + assert execute(project, source) == before == "True\n" + + +def test_shadowed_function_annotation_checks_runtime_namespace(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + target = int + def func(target: str): + def inner(x: target): + pass + return inner + result = target + import typing + print(result is int, typing.get_type_hints(func(str))["x"] is int) + ''')) + before = execute(project, source) + inline = create_inline(project, source, source.read().index("target") + 1) + if deferred_prefix: + code = source.read() + with pytest.raises(RefactoringError, match="deferred annotation"): + inline.get_changes() + assert source.read() == code + assert execute(project, source) == before == "True True\n" + else: + project.do(inline.get_changes()) + assert "result = int" in source.read() + assert "x: target" in source.read() + assert execute(project, source) == before == "True False\n" + + +def test_local_variable_annotation_is_not_evaluated(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + target = int + def func(): + x: target = 1 + return x + print(func()) + ''')) + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "target = int" not in source.read() + assert execute(project, source) == before == "1\n" + + +@pytest.mark.parametrize("annotation", ["holder.x: target = 1", "holder[0]: target = 1", "(result): target = 1"]) +def test_non_simple_annotation_is_not_deferred(project, deferred_prefix, annotation): + if annotation.startswith("(result)"): + setup, access = "", "result" + elif "[0]" in annotation: + setup, access = "holder = [0]\n", "holder[0]" + else: + setup, access = "class Holder: pass\nholder = Holder()\n", "holder.x" + source = module(project, "source", deferred_prefix + "target = int\n" + setup + annotation + f"\nprint({access})\n") + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "target = int" not in source.read() + assert execute(project, source) == before == "1\n" + + +def test_nested_eager_function_annotation_allows_inline(project): + if sys.version_info >= (3, 14): + pytest.skip("default annotations are deferred on Python 3.14+") + source = module(project, "source", '''\ + def outer(): + original = int + target = original + def inner(x: target): + pass + original = str + return inner + import typing + print(typing.get_type_hints(outer())["x"] is int) + ''') + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "target" not in source.read() + assert execute(project, source) == before == "True\n" + + +def test_class_annotation_within_function_is_still_deferred(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + original = int + target = original + def factory(): + class Owner: + x: target + return Owner + original = str + import typing + print(typing.get_type_hints(factory())["x"] is int) + ''')) + before = source.read() + assert execute(project, source) == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + + +def test_class_shadowed_annotation_checks_runtime_namespace(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + target = int + class Owner: + target = str + x: target + result = target + import typing + print(result is int, typing.get_type_hints(Owner)["x"] is str) + ''')) + before = execute(project, source) + inline = create_inline(project, source, source.read().index("target") + 1) + if deferred_prefix: + code = source.read() + with pytest.raises(RefactoringError, match="deferred annotation"): + inline.get_changes() + assert source.read() == code + assert execute(project, source) == before == "True False\n" + else: + project.do(inline.get_changes()) + assert "result = int" in source.read() + assert "target = str" in source.read() + assert execute(project, source) == before == "True True\n" + + +def test_resources_sequence_is_preserved(project, deferred_prefix): + source = module(project, "source", deferred_prefix + "target = int\nresult = target\nprint(result is int)\n") + before = execute(project, source) + resources = (source,) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes(resources=resources)) + assert "target" not in source.read() + assert execute(project, source) == before == "True\n" + + +def test_function_body_shadowing_still_allows_inline(project, deferred_prefix): + source = module(project, "source", deferred_prefix + dedent('''\ + target = int + def func(target: str): + return target + result = target + print(result is int, func(str) is str) + ''')) + before = execute(project, source) + project.do(create_inline(project, source, source.read().index("target") + 1).get_changes()) + assert "result = int" in source.read() + assert "return target" in source.read() + assert execute(project, source) == before == "True True\n" + + +@pytest.mark.parametrize("remove", [False, True]) +@pytest.mark.parametrize("only_current", [False, True]) +def test_only_current_checks_annotation_in_third_module(project, deferred_prefix, remove, only_current): + source = module(project, "source", "target = int\n") + annotated = module(project, "annotated", deferred_prefix + "import source\ndef func(x: source.target):\n pass\n") + entry = module(project, "entry", '''\ + import source + import annotated + result = source.target + import typing + print(result is typing.get_type_hints(annotated.func)["x"]) + ''') + before = {r: r.read() for r in (source, annotated, entry)} + output = execute(project, entry) + inline = create_inline(project, entry, entry.read().rindex("target") + 1) + if remove: + with pytest.raises(RefactoringError, match="deferred annotation"): + inline.get_changes(only_current=only_current, resources=[entry, source]) + assert {r: r.read() for r in before} == before + else: + project.do(inline.get_changes(only_current=only_current, remove=False, resources=[entry, source])) + assert "result = int" in entry.read() + assert source.read() == before[source] + assert annotated.read() == before[annotated] + assert execute(project, entry) == output == "True\n" + + +def test_future_import_alias_preserves_runtime_global_binding(project): + source = module(project, "source", "target = int\n") + entry = module(project, "entry", '''\ + from __future__ import annotations + from source import target as Alias + class Owner: + Alias = str + x: Alias + import typing + print(typing.get_type_hints(Owner)["x"] is int) + ''') + before = {source: source.read(), entry: entry.read()} + assert execute(project, entry) == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, source.read().index("target") + 1).get_changes() + assert {r: r.read() for r in before} == before + assert execute(project, entry) == "True\n" + + +def test_future_qualified_annotation_preserves_runtime_global_binding(project): + source = module(project, "source", "target = int\n") + entry = module(project, "entry", '''\ + from __future__ import annotations + import source + class Holder: + target = str + class Owner: + source = Holder + x: source.target + result = source.target + import typing + print(result is typing.get_type_hints(Owner)["x"]) + ''') + before = {source: source.read(), entry: entry.read()} + assert execute(project, entry) == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, source.read().index("target") + 1).get_changes() + assert {r: r.read() for r in before} == before + assert execute(project, entry) == "True\n" diff --git a/ropetest/refactor/inline_type_parameters_test.py b/ropetest/refactor/inline_type_parameters_test.py new file mode 100644 index 00000000..563b3bf0 --- /dev/null +++ b/ropetest/refactor/inline_type_parameters_test.py @@ -0,0 +1,690 @@ +import subprocess +import sys +from textwrap import dedent + +import pytest + +from rope.base.exceptions import RefactoringError +from rope.base.project import Project +from rope.refactor.inline import create_inline + +pytestmark = pytest.mark.skipif( + sys.version_info < (3, 12), reason="type parameter syntax requires Python 3.12+" +) + + +@pytest.fixture +def project(tmp_path): + project = Project( + str(tmp_path), save_objectdb=False, save_history=False, automatic_soa=False + ) + yield project + project.close() + + +def module(project, name, code): + resource = project.root.create_file(name + ".py") + resource.write(dedent(code)) + return resource + + +def execute(project, resource): + result = subprocess.run( + [sys.executable, resource.real_path], + cwd=project.address, + capture_output=True, + text=True, + timeout=10, + ) + assert result.returncode == 0, result.stderr + return result.stdout + + +@pytest.mark.parametrize( + "declaration,owner", + [ + ("def func[T: target]():", "func"), + ("async def func[T: target]():", "func"), + ("class Owner[T: target]:", "Owner"), + ("def func[T: (target, float)]():", "func"), + ("class Owner[T: (target, float)]:", "Owner"), + ], +) +@pytest.mark.parametrize("remove", [True, False]) +def test_refuses_changing_bound_or_constraint_binding( + project, declaration, owner, remove +): + attribute = "__constraints__" if "float" in declaration else "__bound__" + expected = "(int, float)" if "float" in declaration else "int" + source = module( + project, + "source", + dedent(f"""\ + original = int + target = original + {declaration} + pass + original = str + print({owner}.__type_params__[0].{attribute} == {expected}) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes( + remove=remove + ) + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.skipif( + sys.version_info < (3, 13), reason="type defaults require Python 3.13+" +) +@pytest.mark.parametrize( + "declaration,owner,initial,expected,rebound", + [ + ("def func[T = target]():", "func", "int", "int", "str"), + ("class Owner[T = target]:", "Owner", "int", "int", "str"), + ("def func[**P = target]():", "func", "[int]", "[int]", "[str]"), + ( + "def func[*Ts = *target]():", + "func", + "tuple[int, float]", + "next(iter(tuple[int, float]))", + "tuple[str, float]", + ), + ], +) +def test_refuses_changing_default_binding( + project, declaration, owner, initial, expected, rebound +): + source = module( + project, + "source", + dedent(f"""\ + original = {initial} + target = original + {declaration} + pass + original = {rebound} + print({owner}.__type_params__[0].__default__ == {expected}) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +def test_refuses_deferring_initializer_side_effect(project): + source = module( + project, + "source", + dedent("""\ + events = [] + def make(): + events.append("called") + return int + target = make() + def func[T: target](): + pass + print(events) + print(func.__type_params__[0].__bound__ is int) + print(events) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "['called']\nTrue\n['called']\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize("selection", ["definition", "reference"]) +def test_refuses_removing_binding_outside_selected_write_scope(project, selection): + source = module(project, "source", "target = int\n") + annotation = module( + project, + "annotation", + dedent("""\ + import source + class Owner[T: source.target]: + pass + """), + ) + entry = module( + project, + "entry", + dedent("""\ + import source + import annotation + result = source.target + print(result is int, annotation.Owner.__type_params__[0].__bound__ is int) + """), + ) + files = (source, annotation, entry) + before = [resource.read() for resource in files] + output = execute(project, entry) + assert output == "True True\n" + resource = source if selection == "definition" else entry + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline( + project, resource, resource.read().index("target") + 1 + ).get_changes(only_current=selection == "reference", resources=(entry,)) + assert [resource.read() for resource in files] == before + assert execute(project, entry) == output + + +def test_allows_partial_eager_inline_with_type_parameter_binding_retained(project): + source = module( + project, + "source", + dedent("""\ + original = int + target = original + def func[T: target](): + pass + result = target + original = str + print(result is int, func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True True\n" + project.do( + create_inline( + project, source, before.index("result = target") + len("result = t") + ).get_changes(only_current=True, remove=False) + ) + assert "target = original" in source.read() + assert "result = original" in source.read() + assert execute(project, source) == output + + +def test_allows_eager_reference_in_generic_function_body(project): + source = module( + project, + "source", + dedent("""\ + target = 42 + def func[T: int](): + return target + print(func(), func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "42 True\n" + project.do(create_inline(project, source, before.index("target") + 1).get_changes()) + assert "return 42" in source.read() + assert execute(project, source) == output + + +def test_refuses_removing_captured_local_binding(project): + source = module( + project, + "source", + dedent("""\ + def outer(): + target = int + def func[T: target](): + pass + return func + print(outer().__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +def test_allows_unrelated_shadowed_type_parameter_reference(project): + source = module( + project, + "source", + dedent("""\ + target = 42 + def outer(): + target = int + def func[T: target](): + pass + return func + result = target + print(result, outer().__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "42 True\n" + project.do(create_inline(project, source, before.index("target") + 1).get_changes()) + assert "result = 42" in source.read() + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "definition,owner", + [ + ( + dedent("""\ + def func[T: target](target): + pass + """), + "func", + ), + ( + dedent("""\ + async def func[T: target](target): + pass + """), + "func", + ), + ( + dedent("""\ + def func[T: target](): + target = str + """), + "func", + ), + ( + dedent("""\ + class Owner[T: target]: + target = str + """), + "Owner", + ), + ], +) +def test_refuses_dependency_hidden_by_function_or_class_body( + project, definition, owner +): + source = module( + project, + "source", + dedent("original = int\ntarget = original\n") + definition + dedent(f"""\ + original = str + print({owner}.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "definition,owner", + [ + ( + dedent("""\ + def func[T: target](target): + pass + """), + "Outer.func", + ), + ( + dedent("""\ + class Owner[T: target]: + target = str + """), + "Outer.Owner", + ), + ], +) +def test_refuses_removing_enclosing_class_namespace_binding(project, definition, owner): + code = dedent("class Outer:\n target = int\n") + "".join( + " " + line for line in definition.splitlines(True) + ) + code += f"print({owner}.__type_params__[0].__bound__ is int)\n" + source = module(project, "source", code) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "declaration,expected", + [ + ("def func[target: target]():", "func.__type_params__[0]"), + ("def func[target, T: target]():", "func.__type_params__[0]"), + ], +) +def test_allows_removing_binding_shadowed_by_type_parameter( + project, declaration, expected +): + index = 1 if ", T" in declaration else 0 + source = module( + project, + "source", + dedent(f"""\ + target = int + {declaration} + pass + result = target + print(result is int, func.__type_params__[{index}].__bound__ is {expected}) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True True\n" + project.do( + create_inline( + project, source, before.index("result = target") + len("result = t") + ).get_changes(only_current=True) + ) + assert "result = int" in source.read() + assert execute(project, source) == output + + +@pytest.mark.skipif( + sys.version_info < (3, 13), reason="type defaults require Python 3.13+" +) +@pytest.mark.parametrize( + "definition,owner", + [ + ( + dedent("""\ + def func[T = target](target): + pass + """), + "func", + ), + ( + dedent("""\ + class Owner[T = target]: + target = str + """), + "Owner", + ), + ], +) +def test_refuses_default_dependency_hidden_by_declared_scope( + project, definition, owner +): + source = module( + project, + "source", + dedent("original = int\ntarget = original\n") + definition + dedent(f"""\ + original = str + print({owner}.__type_params__[0].__default__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.skipif( + sys.version_info < (3, 14), + reason="nested annotation scopes in classes require Python 3.14+", +) +@pytest.mark.parametrize( + "expression", + [ + "(lambda: target)()", + "[target for unused in [0]][0]", + "next(target for unused in [0])", + "{unused: target for unused in [0]}[0]", + "{target for unused in [0]}.pop()", + ], +) +def test_refuses_inner_expression_resolving_module_namespace(project, expression): + source = module( + project, + "source", + dedent(f"""\ + original = int + target = original + class Outer: + target = bytes + def func[T: {expression}](target): + pass + original = str + print(Outer.func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "expression", + [ + "(lambda target: target)(int)", + "[target for target in [int]][0]", + "next(target for target in [int])", + "{target: target for target in [int]}[int]", + "{target for target in [int]}.pop()", + ], +) +def test_allows_removing_binding_shadowed_by_inner_expression(project, expression): + source = module( + project, + "source", + dedent(f"""\ + target = 42 + def func[T: {expression}](): + pass + result = target + print(result, func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "42 True\n" + project.do( + create_inline( + project, source, before.index("result = target") + len("result = t") + ).get_changes(only_current=True) + ) + assert "result = 42" in source.read() + assert execute(project, source) == output + + +@pytest.mark.skipif( + sys.version_info < (3, 14), + reason="nested annotation scopes in classes require Python 3.14+", +) +@pytest.mark.parametrize( + "expression", + [ + "(lambda value=target: value)()", + "[value for value in [target]][0]", + ], +) +def test_refuses_inner_expression_outer_class_namespace_dependency(project, expression): + source = module( + project, + "source", + dedent(f"""\ + class Outer: + target = int + def func[T: {expression}](target): + pass + print(Outer.func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "definition,comparison", + [ + ( + dedent("""\ + def outer[target](): + def func[T: target](): + pass + return func.__type_params__[0].__bound__ is outer.__type_params__[0] + """), + "outer()", + ), + ( + dedent("""\ + class Outer[target]: + def func[T: target](): + pass + """), + "Outer.func.__type_params__[0].__bound__ is Outer.__type_params__[0]", + ), + ( + dedent("""\ + class Outer[target]: + target = bytes + def method(self): + def func[T: target](): + pass + return func.__type_params__[0].__bound__ is Outer.__type_params__[0] + """), + "Outer().method()", + ), + ], +) +def test_allows_removing_binding_shadowed_by_enclosing_type_parameter( + project, definition, comparison +): + source = module( + project, + "source", + "target = 42\n" + definition + dedent(f"""\ + result = target + print(result, {comparison}) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "42 True\n" + project.do( + create_inline( + project, source, before.index("result = target") + len("result = t") + ).get_changes(only_current=True) + ) + assert "result = 42" in source.read() + assert execute(project, source) == output + + +def test_refuses_local_binding_inside_enclosing_generic_function(project): + source = module( + project, + "source", + dedent("""\ + def outer[target](): + def inner(): + target = int + def func[T: target](): + pass + return func + return inner() + print(outer().__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target = int") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.parametrize( + "initializer,assignment", + [ + ("Box()", "target.attr"), + ("{}", "target[0]"), + ], +) +def test_refuses_object_dependency_in_comprehension_assignment_target( + project, initializer, assignment +): + source = module( + project, + "source", + dedent(f"""\ + class Box: + pass + target = {initializer} + def func[T: [int for {assignment} in [int]][0]](target): + pass + print(func.__type_params__[0].__bound__ is int) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "True\n" + with pytest.raises(RefactoringError, match="deferred annotation"): + create_inline(project, source, before.index("target") + 1).get_changes() + assert source.read() == before + assert execute(project, source) == output + + +@pytest.mark.skipif( + sys.version_info < (3, 14), + reason="nested annotation scopes in classes require Python 3.14+", +) +@pytest.mark.parametrize( + "expression", + [ + "(lambda: target)()", + "[target for unused in [0]][0]", + ], +) +def test_allows_inner_expression_capturing_enclosing_class_type_parameter( + project, expression +): + source = module( + project, + "source", + dedent(f"""\ + target = 42 + class Owner[target]: + target = bytes + def func[T: {expression}](): + pass + result = target + print(result, Owner.func.__type_params__[0].__bound__ is Owner.__type_params__[0]) + """), + ) + before = source.read() + output = execute(project, source) + assert output == "42 True\n" + project.do( + create_inline( + project, source, before.index("result = target") + len("result = t") + ).get_changes(only_current=True) + ) + assert "result = 42" in source.read() + assert execute(project, source) == output