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