Skip to content
227 changes: 195 additions & 32 deletions src/_griffe/expressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,8 @@
import sys
from dataclasses import dataclass
from dataclasses import fields as getfields
from enum import IntEnum, auto
from functools import partial
from itertools import zip_longest
from typing import TYPE_CHECKING, Any, Callable

from _griffe.agents.nodes.parameters import get_parameters
Expand All @@ -26,14 +26,82 @@
from _griffe.models import Class, Function, Module


def _yield(element: str | Expr | tuple[str | Expr, ...], *, flat: bool = True) -> Iterator[str | Expr]:
if isinstance(element, str):
yield element
class _OperatorPrecedence(IntEnum):
# Adapted from:
#
# - https://docs.python.org/3/reference/expressions.html#operator-precedence
# - https://github.com/python/cpython/blob/main/Lib/_ast_unparse.py
# - https://github.com/astral-sh/ruff/blob/6abafcb56575454f2caeaa174efcb9fd0a8362b1/crates/ruff_python_ast/src/operator_precedence.rs

# The enum members are declared in ascending order of precedence.

# A virtual precedence level for contexts that provide their own grouping, like list brackets or
# function call parentheses. This ensures parentheses will never be added for the direct children of these nodes.
# NOTE: `ruff_python_formatter::expression::parentheses`'s state machine would be more robust
# but would introduce significant complexity.
NONE = auto()

YIELD = auto() # `yield`, `yield from`
ASSIGN = auto() # `target := expr`
STARRED = auto() # `*expr` (omitted by Python docs, see ruff impl)
LAMBDA = auto()
IF_ELSE = auto() # `expr if cond else expr`
OR = auto()
AND = auto()
NOT = auto()
COMP_MEMB_ID = auto() # `<`, `<=`, `>`, `>=`, `!=`, `==`, `in`, `not in`, `is`, `is not`
BIT_OR = auto() # `|`
BIT_XOR = auto() # `^`
BIT_AND = auto() # `&`
LEFT_RIGHT_SHIFT = auto() # `<<`, `>>`
ADD_SUB = auto() # `+`, `-`
MUL_DIV_REMAIN = auto() # `*`, `@`, `/`, `//`, `%`
POS_NEG_BIT_NOT = auto() # `+x`, `-x`, `~x`
EXPONENT = auto() # `**`
AWAIT = auto()
CALL_ATTRIBUTE = auto() # `x[index]`, `x[index:index]`, `x(arguments...)`, `x.attribute`
ATOMIC = auto() # `(expressions...)`, `[expressions...]`, `{key: value...}`, `{expressions...}`


def _yield(
element: str | Expr | tuple[str | Expr, ...],
*,
flat: bool = True,
is_left: bool = False,
outer_precedence: _OperatorPrecedence = _OperatorPrecedence.ATOMIC,
) -> Iterator[str | Expr]:
if isinstance(element, Expr):
element_precedence = _get_precedence(element)
needs_parens = False
# Lower inner precedence, e.g. `(a + b) * c`, `+(10) < *(11)`.
if element_precedence < outer_precedence:
needs_parens = True
elif element_precedence == outer_precedence:
# Right-association, e.g. parenthesize left-hand side in `(a ** b) ** c`, (a if b else c) if d else e
is_right_assoc = isinstance(element, ExprIfExp) or (
isinstance(element, ExprBinOp) and element.operator == "**"
)
if is_right_assoc:
if is_left:
needs_parens = True
# Left-association, e.g. parenthesize right-hand side in `a - (b - c)`.
elif isinstance(element, (ExprBinOp, ExprBoolOp)) and not is_left:
needs_parens = True

if needs_parens:
yield "("
if flat:
yield from element.iterate(flat=True)
else:
yield element
yield ")"
elif flat:
yield from element.iterate(flat=True)
else:
yield element
elif isinstance(element, tuple):
for elem in element:
yield from _yield(elem, flat=flat)
elif flat:
yield from element.iterate(flat=True)
yield from _yield(elem, flat=flat, outer_precedence=outer_precedence, is_left=is_left)
else:
yield element

Expand All @@ -44,14 +112,21 @@ def _join(
*,
flat: bool = True,
) -> Iterator[str | Expr]:
"""Apply a separator between elements.

The caller is assumed to provide their own grouping
(e.g. lists, tuples, slice) and will prevent parentheses from being added.
"""
it = iter(elements)
try:
yield from _yield(next(it), flat=flat)
# Since we are in a sequence, don't parenthesize items.
# Avoids [a + b, c + d] being serialized as [(a + b), (c + d)]
yield from _yield(next(it), flat=flat, outer_precedence=_OperatorPrecedence.NONE)
except StopIteration:
return
for element in it:
yield from _yield(joint, flat=flat)
yield from _yield(element, flat=flat)
yield from _yield(joint, flat=flat, outer_precedence=_OperatorPrecedence.NONE)
yield from _yield(element, flat=flat, outer_precedence=_OperatorPrecedence.NONE)


def _field_as_dict(
Expand Down Expand Up @@ -184,7 +259,11 @@ class ExprAttribute(Expr):
"""The different parts of the dotted chain."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _join(self.values, ".", flat=flat)
precedence = _get_precedence(self)
yield from _yield(self.values[0], flat=flat, outer_precedence=precedence, is_left=True)
for value in self.values[1:]:
yield "."
yield from _yield(value, flat=flat, outer_precedence=precedence)

def append(self, value: ExprName) -> None:
"""Append a name to this attribute.
Expand Down Expand Up @@ -232,9 +311,14 @@ class ExprBinOp(Expr):
"""Right part."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _yield(self.left, flat=flat)
precedence = _get_precedence(self)
right_precedence = precedence
if self.operator == "**" and isinstance(self.right, ExprUnaryOp):
# Unary operators on the right have higher precedence, e.g. `a ** -b`.
right_precedence = _OperatorPrecedence(precedence - 1)
yield from _yield(self.left, flat=flat, outer_precedence=precedence, is_left=True)
yield f" {self.operator} "
yield from _yield(self.right, flat=flat)
yield from _yield(self.right, flat=flat, outer_precedence=right_precedence, is_left=False)


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand All @@ -248,7 +332,12 @@ class ExprBoolOp(Expr):
"""Operands."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _join(self.values, f" {self.operator} ", flat=flat)
precedence = _get_precedence(self)
it = iter(self.values)
yield from _yield(next(it), flat=flat, outer_precedence=precedence, is_left=True)
for value in it:
yield f" {self.operator} "
yield from _yield(value, flat=flat, outer_precedence=precedence, is_left=False)


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand All @@ -267,7 +356,7 @@ def canonical_path(self) -> str:
return self.function.canonical_path

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _yield(self.function, flat=flat)
yield from _yield(self.function, flat=flat, outer_precedence=_get_precedence(self))
yield "("
yield from _join(self.arguments, ", ", flat=flat)
yield ")"
Expand All @@ -286,9 +375,11 @@ class ExprCompare(Expr):
"""Things compared."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _yield(self.left, flat=flat)
yield " "
yield from _join(zip_longest(self.operators, [], self.comparators, fillvalue=" "), " ", flat=flat)
precedence = _get_precedence(self)
yield from _yield(self.left, flat=flat, outer_precedence=precedence, is_left=True)
for op, comp in zip(self.operators, self.comparators):
yield f" {op} "
yield from _yield(comp, flat=flat, outer_precedence=precedence)
Comment thread
abc8747 marked this conversation as resolved.


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand Down Expand Up @@ -396,7 +487,8 @@ class ExprFormatted(Expr):

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield "{"
yield from _yield(self.value, flat=flat)
# Prevent parentheses from being added, avoiding `{(1 + 1)}`
yield from _yield(self.value, flat=flat, outer_precedence=_OperatorPrecedence.NONE)
yield "}"


Expand Down Expand Up @@ -429,11 +521,21 @@ class ExprIfExp(Expr):
"""Other expression."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _yield(self.body, flat=flat)
precedence = _get_precedence(self)
yield from _yield(self.body, flat=flat, outer_precedence=precedence, is_left=True)
yield " if "
yield from _yield(self.test, flat=flat)
# If the test itself is another if/else, its precedence is the same, which will not give
# a parenthesis: force it.
test_outer_precedence = _OperatorPrecedence(precedence + 1)
yield from _yield(self.test, flat=flat, outer_precedence=test_outer_precedence)
yield " else "
yield from _yield(self.orelse, flat=flat)
# If/else is right associative. For example, a nested if/else
# `a if b else c if d else e` is effectively `a if b else (c if d else e)`, so produce a
# flattened version without parentheses.
if isinstance(self.orelse, ExprIfExp):
yield from self.orelse.iterate(flat=flat)
else:
yield from _yield(self.orelse, flat=flat, outer_precedence=precedence, is_left=False)


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand Down Expand Up @@ -557,7 +659,8 @@ def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
if index < length:
yield ", "
yield ": "
yield from _yield(self.body, flat=flat)
# Body of lambda should not have parentheses, avoiding `lambda: a.b`
yield from _yield(self.body, flat=flat, outer_precedence=_OperatorPrecedence.NONE)


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand Down Expand Up @@ -685,11 +788,9 @@ class ExprNamedExpr(Expr):
"""Value."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield "("
yield from _yield(self.target, flat=flat)
yield " := "
yield from _yield(self.value, flat=flat)
yield ")"


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand Down Expand Up @@ -773,9 +874,10 @@ class ExprSubscript(Expr):
"""Slice part."""

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield from _yield(self.left, flat=flat)
yield from _yield(self.left, flat=flat, outer_precedence=_get_precedence(self))
yield "["
yield from _yield(self.slice, flat=flat)
# Prevent parentheses from being added, avoiding `a[(b)]`
yield from _yield(self.slice, flat=flat, outer_precedence=_OperatorPrecedence.NONE)
yield "]"

@property
Expand Down Expand Up @@ -825,7 +927,9 @@ class ExprUnaryOp(Expr):

def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
yield self.operator
yield from _yield(self.value, flat=flat)
if self.operator == "not":
yield " "
yield from _yield(self.value, flat=flat, outer_precedence=_get_precedence(self))


# YORE: EOL 3.9: Replace `**_dataclass_opts` with `slots=True` within line.
Expand Down Expand Up @@ -858,7 +962,7 @@ def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:

_unary_op_map = {
ast.Invert: "~",
ast.Not: "not ",
ast.Not: "not",
ast.UAdd: "+",
ast.USub: "-",
}
Expand Down Expand Up @@ -897,6 +1001,65 @@ def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]:
ast.NotIn: "not in",
}

# TODO: Support `ast.Await`.
_precedence_map = {
# Literals and names.
ExprName: lambda _: _OperatorPrecedence.ATOMIC,
ExprConstant: lambda _: _OperatorPrecedence.ATOMIC,
ExprJoinedStr: lambda _: _OperatorPrecedence.ATOMIC,
ExprFormatted: lambda _: _OperatorPrecedence.ATOMIC,
# Container displays.
ExprList: lambda _: _OperatorPrecedence.ATOMIC,
ExprTuple: lambda _: _OperatorPrecedence.ATOMIC,
ExprSet: lambda _: _OperatorPrecedence.ATOMIC,
ExprDict: lambda _: _OperatorPrecedence.ATOMIC,
# Comprehensions are self-contained units that produce a container.
ExprListComp: lambda _: _OperatorPrecedence.ATOMIC,
ExprSetComp: lambda _: _OperatorPrecedence.ATOMIC,
ExprDictComp: lambda _: _OperatorPrecedence.ATOMIC,
ExprAttribute: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE,
ExprSubscript: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE,
ExprCall: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE,
ExprUnaryOp: lambda e: {"not": _OperatorPrecedence.NOT}.get(e.operator, _OperatorPrecedence.POS_NEG_BIT_NOT),
ExprBinOp: lambda e: {
"**": _OperatorPrecedence.EXPONENT,
"*": _OperatorPrecedence.MUL_DIV_REMAIN,
"@": _OperatorPrecedence.MUL_DIV_REMAIN,
"/": _OperatorPrecedence.MUL_DIV_REMAIN,
"//": _OperatorPrecedence.MUL_DIV_REMAIN,
"%": _OperatorPrecedence.MUL_DIV_REMAIN,
"+": _OperatorPrecedence.ADD_SUB,
"-": _OperatorPrecedence.ADD_SUB,
"<<": _OperatorPrecedence.LEFT_RIGHT_SHIFT,
">>": _OperatorPrecedence.LEFT_RIGHT_SHIFT,
"&": _OperatorPrecedence.BIT_AND,
"^": _OperatorPrecedence.BIT_XOR,
"|": _OperatorPrecedence.BIT_OR,
}[e.operator],
ExprBoolOp: lambda e: {"and": _OperatorPrecedence.AND, "or": _OperatorPrecedence.OR}[e.operator],
ExprCompare: lambda _: _OperatorPrecedence.COMP_MEMB_ID,
ExprIfExp: lambda _: _OperatorPrecedence.IF_ELSE,
ExprNamedExpr: lambda _: _OperatorPrecedence.ASSIGN,
ExprLambda: lambda _: _OperatorPrecedence.LAMBDA,
# NOTE: Ruff categorizes as atomic, but `(a for a in b).c` implies its less than `CALL_ATTRIBUTE`.
ExprGeneratorExp: lambda _: _OperatorPrecedence.LAMBDA,
ExprVarPositional: lambda _: _OperatorPrecedence.STARRED,
ExprVarKeyword: lambda _: _OperatorPrecedence.STARRED,
ExprYield: lambda _: _OperatorPrecedence.YIELD,
ExprYieldFrom: lambda _: _OperatorPrecedence.YIELD,
# These are not standalone, they appear in specific contexts where precendence is not a concern.
# NOTE: `for ... in ... if` part, not the whole `[...]`.
ExprComprehension: lambda _: _OperatorPrecedence.NONE,
ExprExtSlice: lambda _: _OperatorPrecedence.NONE,
ExprKeyword: lambda _: _OperatorPrecedence.NONE,
ExprParameter: lambda _: _OperatorPrecedence.NONE,
ExprSlice: lambda _: _OperatorPrecedence.NONE,
}


def _get_precedence(expr: Expr) -> _OperatorPrecedence:
return _precedence_map.get(type(expr), lambda _: _OperatorPrecedence.NONE)(expr)


def _build_attribute(node: ast.Attribute, parent: Module | Class, **kwargs: Any) -> Expr:
left = _build(node.value, parent, **kwargs)
Expand Down Expand Up @@ -1113,7 +1276,7 @@ def _build_subscript(
"typing_extensions.Literal",
}:
literal_strings = True
slice = _build(
slice_expr = _build(
node.slice,
parent,
parse_strings=True,
Expand All @@ -1122,8 +1285,8 @@ def _build_subscript(
**kwargs,
)
else:
slice = _build(node.slice, parent, in_subscript=True, **kwargs)
return ExprSubscript(left, slice)
slice_expr = _build(node.slice, parent, in_subscript=True, **kwargs)
return ExprSubscript(left, slice_expr)


def _build_tuple(
Expand Down
Loading
Loading