diff --git a/rewrite-python/rewrite/src/rewrite/python/tree.py b/rewrite-python/rewrite/src/rewrite/python/tree.py index 9b18ea8801..9337a0f7ba 100644 --- a/rewrite-python/rewrite/src/rewrite/python/tree.py +++ b/rewrite-python/rewrite/src/rewrite/python/tree.py @@ -19,6 +19,20 @@ ) from rewrite.python.support_types import Py, P + +def _delegating_prefix_and_markers(child: J, child_key: str, kwargs: dict) -> dict: + """Redirect prefix/markers onto the wrapped child. + + `ExpressionStatement` and `StatementExpression` store neither of their own, + exposing the child's instead, so a plain `replace` would silently drop them. + """ + if 'prefix' in kwargs: + kwargs[child_key] = child.replace(prefix=kwargs.pop('prefix')) + if 'markers' in kwargs: + kwargs[child_key] = kwargs.get(child_key, child).replace(markers=kwargs.pop('markers')) + return kwargs + + # noinspection PyShadowingBuiltins,PyShadowingNames,DuplicatedCode @dataclass(frozen=True, eq=False, slots=True) class Async(Py, Statement): @@ -505,6 +519,9 @@ def expression(self) -> Expression: return self._expression + def replace(self, **kwargs) -> 'ExpressionStatement': + return replace_if_changed(self, **_delegating_prefix_and_markers(self._expression, 'expression', kwargs)) + def accept_python(self, v: PythonVisitor[P], p: P) -> J: return v.visit_expression_statement(self, p) @@ -560,15 +577,7 @@ def statement(self) -> Statement: def replace(self, **kwargs) -> 'StatementExpression': - """Replace fields, handling delegated prefix/markers specially.""" - # Handle delegated properties by modifying the inner statement - if 'prefix' in kwargs: - new_statement = self._statement.replace(prefix=kwargs.pop('prefix')) # Statement base class doesn't have replace - kwargs['statement'] = new_statement - if 'markers' in kwargs: - new_statement = kwargs.get('statement', self._statement).replace(markers=kwargs.pop('markers')) - kwargs['statement'] = new_statement - return replace_if_changed(self, **kwargs) + return replace_if_changed(self, **_delegating_prefix_and_markers(self._statement, 'statement', kwargs)) def accept_python(self, v: PythonVisitor[P], p: P) -> J: return v.visit_statement_expression(self, p) diff --git a/rewrite-python/rewrite/src/rewrite/python/visitor.py b/rewrite-python/rewrite/src/rewrite/python/visitor.py index 4557e116d7..bf871521a9 100644 --- a/rewrite-python/rewrite/src/rewrite/python/visitor.py +++ b/rewrite-python/rewrite/src/rewrite/python/visitor.py @@ -17,6 +17,7 @@ from __future__ import annotations from typing import TYPE_CHECKING, Any, TypeVar, cast, Optional +from uuid import UUID from rewrite.java import tree as j from rewrite.java.support_types import ( @@ -26,6 +27,7 @@ ) from rewrite.java.visitor import JavaVisitor from rewrite.python.support_types import Py +from rewrite.python.tree import ExpressionStatement, StatementExpression from rewrite.tree import SourceFile from rewrite.utils import list_map from rewrite.visitor import TreeVisitor @@ -43,7 +45,6 @@ DictLiteral, ErrorFrom, ExceptionType, - ExpressionStatement, ExpressionTypeTree, FormattedString, KeyValue, @@ -55,7 +56,6 @@ Slice, SpecialParameter, Star, - StatementExpression, TrailingElseWrapper, TypeAlias, TypeHint, @@ -69,6 +69,23 @@ T = TypeVar("T") +def _expression_statement_kinds(node: J) -> tuple[bool, bool]: + """Which of the two roles `node` already fills, deciding the wrapper it still needs.""" + return isinstance(node, Expression), isinstance(node, Statement) + + +def _minimal_expression_statement_wrapper(wrapper_id: UUID, child: J) -> J: + """The least wrapping that makes `child` usable as both an expression and a statement. + + Both wrappers are themselves expression and statement, so an already wrapped + child needs nothing further and cannot become doubly wrapped. + """ + is_expression, is_statement = _expression_statement_kinds(child) + if is_expression: + return child if is_statement else ExpressionStatement(wrapper_id, child) + return StatementExpression(wrapper_id, child) if is_statement else child + + class PythonVisitor(JavaVisitor[P]): """ Base visitor for Python LST nodes. @@ -310,7 +327,7 @@ def visit_exception_type(self, exc_type: ExceptionType, p: P) -> J: ) return exc_type - def visit_expression_statement(self, expr_stmt: ExpressionStatement, p: P) -> J: + def visit_expression_statement(self, expr_stmt: ExpressionStatement, p: P) -> Optional[J]: """Visit an expression used as a statement.""" temp_stmt = cast(Statement, self.visit_statement(expr_stmt, p)) if not isinstance(temp_stmt, type(expr_stmt)): @@ -320,9 +337,12 @@ def visit_expression_statement(self, expr_stmt: ExpressionStatement, p: P) -> J: if not isinstance(temp_expr, type(expr_stmt)): return temp_expr expr_stmt = temp_expr - expr_stmt = expr_stmt.replace( - expression=self.visit_and_cast(expr_stmt.expression, Expression, p) - ) + expression = self.visit_and_cast(expr_stmt.expression, Expression, p) + if expression is None: + return None + if _expression_statement_kinds(expression) != _expression_statement_kinds(expr_stmt.expression): + return _minimal_expression_statement_wrapper(expr_stmt.id, expression) + expr_stmt = expr_stmt.replace(expression=expression) return expr_stmt def visit_expression_type_tree(self, expr_tree: ExpressionTypeTree, p: P) -> J: @@ -536,7 +556,7 @@ def visit_star(self, star: Star, p: P) -> J: ) return star - def visit_statement_expression(self, stmt_expr: StatementExpression, p: P) -> J: + def visit_statement_expression(self, stmt_expr: StatementExpression, p: P) -> Optional[J]: """Visit a statement used as an expression.""" temp_stmt = cast(Statement, self.visit_statement(stmt_expr, p)) if not isinstance(temp_stmt, type(stmt_expr)): @@ -546,9 +566,12 @@ def visit_statement_expression(self, stmt_expr: StatementExpression, p: P) -> J: if not isinstance(temp_expr, type(stmt_expr)): return temp_expr stmt_expr = temp_expr - stmt_expr = stmt_expr.replace( - statement=self.visit_and_cast(stmt_expr.statement, Statement, p) - ) + statement = self.visit_and_cast(stmt_expr.statement, Statement, p) + if statement is None: + return None + if _expression_statement_kinds(statement) != _expression_statement_kinds(stmt_expr.statement): + return _minimal_expression_statement_wrapper(stmt_expr.id, statement) + stmt_expr = stmt_expr.replace(statement=statement) return stmt_expr def visit_trailing_else_wrapper(self, wrapper: TrailingElseWrapper, p: P) -> J: diff --git a/rewrite-python/rewrite/tests/python/test_wrapper_normalization.py b/rewrite-python/rewrite/tests/python/test_wrapper_normalization.py new file mode 100644 index 0000000000..8ad0c093a6 --- /dev/null +++ b/rewrite-python/rewrite/tests/python/test_wrapper_normalization.py @@ -0,0 +1,262 @@ +# Copyright 2026 the original author or authors. +#

+# Licensed under the Moderne Source Available License (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +#

+# https://docs.moderne.io/licensing/moderne-source-available-license +#

+# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for ExpressionStatement/StatementExpression wrapper normalization. + +Regression tests for https://github.com/openrewrite/rewrite/issues/8322: +when a visitor replaces the child of one of these wrappers with a node of a +different kind, the base visitor must re-derive the wrapping instead of +silently constructing an invalid tree (which only fails later, e.g. with a +ClassCastException on RPC receive in Java). +""" + +from typing import Optional + +from rewrite import ExecutionContext, InMemoryExecutionContext, Markers, random_id +from rewrite.java import J +from rewrite.java.support_types import Expression, JLeftPadded, Space, Statement +from rewrite.java.tree import Assignment, Identifier, Literal, Yield +from rewrite.python.tree import Await, ExpressionStatement, StatementExpression, YieldFrom +from rewrite.python.visitor import PythonVisitor +from rewrite.test import RecipeSpec, from_visitor, python + + +def _assert_wrappers_well_formed(source_file): + """Every wrapper must hold a child of the kind it exists to adapt.""" + + class WrapperWalker(PythonVisitor[ExecutionContext]): + def pre_visit(self, tree, p): + if isinstance(tree, StatementExpression): + child = tree.statement + assert isinstance(child, Statement), \ + f"StatementExpression wraps non-Statement {type(child).__name__}" + elif isinstance(tree, ExpressionStatement): + child = tree.expression + assert isinstance(child, Expression), \ + f"ExpressionStatement wraps non-Expression {type(child).__name__}" + else: + return tree + assert not isinstance(child, (StatementExpression, ExpressionStatement)), \ + f"{type(tree).__name__} redundantly wraps {type(child).__name__}" + return tree + + WrapperWalker().visit(source_file, InMemoryExecutionContext()) + + +class _YieldFromToAwaitVisitor(PythonVisitor[ExecutionContext]): + """Mimics recipes that migrate `yield from expr` to `await expr`. + + `Py.Await` is an `Expression` but not a `Statement`, so the enclosing + `Py.StatementExpression` produced by the parser for a bare `yield from` + statement can no longer hold it. + """ + + def visit_yield(self, yield_stmt: Yield, p: ExecutionContext) -> Optional[J]: + yield_stmt = super().visit_yield(yield_stmt, p) + if isinstance(yield_stmt, Yield) and isinstance(yield_stmt.value, YieldFrom): + yield_from = yield_stmt.value + return Await( + random_id(), + yield_stmt.prefix, + yield_stmt.markers, + yield_from.expression, + yield_from.type, + ) + return yield_stmt + + +def test_statement_expression_rewrapped_when_child_becomes_expression(): + """A bare `yield from` statement whose Yield is replaced by Await must not + leave an Await inside a StatementExpression.""" + RecipeSpec(recipe=from_visitor(_YieldFromToAwaitVisitor())).rewrite_run( + python( + """\ + def coro(): + yield from task() + """, + """\ + def coro(): + await task() + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +def test_statement_expression_in_expression_position_rewrapped(): + """`x = yield from expr` puts the StatementExpression in expression + position; the replacement must stay a valid Expression there too.""" + RecipeSpec(recipe=from_visitor(_YieldFromToAwaitVisitor())).rewrite_run( + python( + """\ + def coro(): + x = yield from task() + """, + """\ + def coro(): + x = await task() + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +class _AwaitToYieldFromVisitor(PythonVisitor[ExecutionContext]): + """The mirror case: replaces `await expr` with `yield from expr`. + + `J.Yield` is a `Statement` but not an `Expression`, so the enclosing + `Py.ExpressionStatement` produced by the parser for a bare `await` + statement can no longer hold it. + """ + + def visit_await(self, await_: Await, p: ExecutionContext) -> Optional[J]: + await_ = super().visit_await(await_, p) + if isinstance(await_, Await): + return Yield( + random_id(), + await_.prefix, + await_.markers, + False, + YieldFrom( + random_id(), + Space.SINGLE_SPACE, + Markers.EMPTY, + await_.expression, + await_.type, + ), + ) + return await_ + + +def test_expression_statement_rewrapped_when_child_becomes_statement(): + """A bare `await` statement whose Await is replaced by Yield must not + leave a Yield inside an ExpressionStatement.""" + RecipeSpec(recipe=from_visitor(_AwaitToYieldFromVisitor())).rewrite_run( + python( + """\ + async def coro(): + await task() + """, + """\ + async def coro(): + yield from task() + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +class _YieldToAssignmentVisitor(PythonVisitor[ExecutionContext]): + """Replaces `yield ` with `a = `. + + `J.Assignment` is both an `Expression` and a `Statement`, so it needs no + wrapper at all. Left inside the `Py.StatementExpression`, the printer emits + the walrus form `a := 1`, which is not valid in statement position. + """ + + def visit_yield(self, yield_stmt: Yield, p: ExecutionContext) -> Optional[J]: + yield_stmt = super().visit_yield(yield_stmt, p) + if isinstance(yield_stmt, Yield) and isinstance(yield_stmt.value, Literal): + return Assignment( + random_id(), + yield_stmt.prefix, + yield_stmt.markers, + Identifier(random_id(), Space.EMPTY, Markers.EMPTY, [], 'a', None, None), + JLeftPadded(Space.SINGLE_SPACE, yield_stmt.value.replace(prefix=Space.SINGLE_SPACE), Markers.EMPTY), + None, + ) + return yield_stmt + + +def test_dual_kind_replacement_drops_the_wrapper(): + """A replacement that is already both Expression and Statement needs no + wrapper; keeping one makes the printer emit `a := 1`.""" + RecipeSpec(recipe=from_visitor(_YieldToAssignmentVisitor())).rewrite_run( + python( + """\ + def gen(): + yield 1 + """, + """\ + def gen(): + a = 1 + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +class _PreWrappedAwaitVisitor(_YieldFromToAwaitVisitor): + """Returns an already-wrapped replacement, as a recipe reasonably might.""" + + def visit_yield(self, yield_stmt: Yield, p: ExecutionContext) -> Optional[J]: + replacement = super().visit_yield(yield_stmt, p) + if isinstance(replacement, Await): + return ExpressionStatement(random_id(), replacement) + return replacement + + +def test_already_wrapped_replacement_is_not_double_wrapped(): + """An already-wrapped replacement must not gain a second wrapper.""" + RecipeSpec(recipe=from_visitor(_PreWrappedAwaitVisitor())).rewrite_run( + python( + """\ + def coro(): + yield from task() + """, + """\ + def coro(): + await task() + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +class _DeleteYieldVisitor(PythonVisitor[ExecutionContext]): + """Deletes the wrapped child, which must take the wrapper with it.""" + + def visit_yield(self, yield_stmt: Yield, p: ExecutionContext) -> Optional[J]: + return None + + +def test_deleted_child_deletes_the_wrapper(): + """Returning None for the wrapped child removes the whole statement rather + than leaving a wrapper holding None.""" + RecipeSpec(recipe=from_visitor(_DeleteYieldVisitor())).rewrite_run( + python( + """\ + def gen(): + yield 1 + print(2) + """, + """\ + def gen(): + print(2) + """, + after_recipe=_assert_wrappers_well_formed, + ) + ) + + +def test_plain_yield_untouched(): + """Plain `yield` (no `yield from`) keeps its StatementExpression wrapper.""" + RecipeSpec(recipe=from_visitor(_YieldFromToAwaitVisitor())).rewrite_run( + python( + """\ + def gen(): + yield 1 + """ + ) + )