Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 18 additions & 9 deletions rewrite-python/rewrite/src/rewrite/python/tree.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)
Expand Down
43 changes: 33 additions & 10 deletions rewrite-python/rewrite/src/rewrite/python/visitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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
Expand All @@ -43,7 +45,6 @@
DictLiteral,
ErrorFrom,
ExceptionType,
ExpressionStatement,
ExpressionTypeTree,
FormattedString,
KeyValue,
Expand All @@ -55,7 +56,6 @@
Slice,
SpecialParameter,
Star,
StatementExpression,
TrailingElseWrapper,
TypeAlias,
TypeHint,
Expand All @@ -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.
Expand Down Expand Up @@ -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)):
Expand All @@ -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:
Expand Down Expand Up @@ -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)):
Expand All @@ -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:
Expand Down
Loading