From 91ad44c0c8d829fe6086569f2fc5f1cccb40aded Mon Sep 17 00:00:00 2001 From: Vincent Gao Date: Thu, 16 Jul 2026 23:43:02 +0200 Subject: [PATCH] feat: Support Await expressions --- packages/griffelib/src/griffe/__init__.py | 2 ++ .../src/griffe/_internal/expressions.py | 19 ++++++++++++- packages/griffelib/tests/test_expressions.py | 27 ++++++++++++++++++- 3 files changed, 46 insertions(+), 2 deletions(-) diff --git a/packages/griffelib/src/griffe/__init__.py b/packages/griffelib/src/griffe/__init__.py index 3d13f377b..a559e0fdf 100644 --- a/packages/griffelib/src/griffe/__init__.py +++ b/packages/griffelib/src/griffe/__init__.py @@ -277,6 +277,7 @@ from griffe._internal.expressions import ( Expr, ExprAttribute, + ExprAwait, ExprBinOp, ExprBoolOp, ExprCall, @@ -433,6 +434,7 @@ "ExplanationStyle", "Expr", "ExprAttribute", + "ExprAwait", "ExprBinOp", "ExprBoolOp", "ExprCall", diff --git a/packages/griffelib/src/griffe/_internal/expressions.py b/packages/griffelib/src/griffe/_internal/expressions.py index 73082b2f6..4423c0526 100644 --- a/packages/griffelib/src/griffe/_internal/expressions.py +++ b/packages/griffelib/src/griffe/_internal/expressions.py @@ -1044,6 +1044,18 @@ def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]: yield from _yield(self.value, flat=flat, outer_precedence=_get_precedence(self)) +@dataclass(eq=True, slots=True) +class ExprAwait(Expr): + """Await expressions like `await call()`.""" + + value: str | Expr + """Awaited value.""" + + def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]: + yield "await " + yield from _yield(self.value, flat=flat, outer_precedence=_OperatorPrecedence.CALL_ATTRIBUTE) + + @dataclass(eq=True, slots=True) class ExprYield(Expr): """Yield statements like `yield a`.""" @@ -1111,7 +1123,6 @@ 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, @@ -1130,6 +1141,7 @@ def iterate(self, *, flat: bool = True) -> Iterator[str | Expr]: ExprAttribute: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE, ExprSubscript: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE, ExprCall: lambda _: _OperatorPrecedence.CALL_ATTRIBUTE, + ExprAwait: lambda _: _OperatorPrecedence.AWAIT, ExprUnaryOp: lambda e: {"not": _OperatorPrecedence.NOT}.get(e.operator, _OperatorPrecedence.POS_NEG_BIT_NOT), ExprBinOp: lambda e: { "**": _OperatorPrecedence.EXPONENT, @@ -1183,6 +1195,10 @@ def _build_attribute(node: ast.Attribute, parent: Module | Class, **kwargs: Any) return ExprAttribute([left, ExprName(node.attr)]) +def _build_await(node: ast.Await, parent: Module | Class, **kwargs: Any) -> Expr: + return ExprAwait(_build(node.value, parent, **kwargs)) + + def _build_binop(node: ast.BinOp, parent: Module | Class, **kwargs: Any) -> Expr: return ExprBinOp( _build(node.left, parent, **kwargs), @@ -1433,6 +1449,7 @@ def __call__(self, node: Any, parent: Module | Class, **kwargs: Any) -> Expr: .. _node_map: dict[type, _BuildCallable] = { ast.Attribute: _build_attribute, + ast.Await: _build_await, ast.BinOp: _build_binop, ast.BoolOp: _build_boolop, ast.Call: _build_call, diff --git a/packages/griffelib/tests/test_expressions.py b/packages/griffelib/tests/test_expressions.py index c18cf52ba..76fae5cb7 100644 --- a/packages/griffelib/tests/test_expressions.py +++ b/packages/griffelib/tests/test_expressions.py @@ -7,7 +7,7 @@ import pytest -from griffe import Module, Parser, get_expression, temporary_visited_module +from griffe import ExprAwait, Module, Parser, get_expression, temporary_visited_module from tests.test_nodes import syntax_examples @@ -88,6 +88,31 @@ def test_expressions(code: str) -> None: assert str(expression) == code +@pytest.mark.parametrize( + ("source", "expected"), + [ + ("await dependency()", "await dependency()"), + ("(await dependency()).attr", "(await dependency()).attr"), + ("(await dependency())()", "(await dependency())()"), + ("await dependency() ** exponent", "await dependency() ** exponent"), + ("-await dependency()", "-await dependency()"), + ("await (dependency() ** exponent)", "await (dependency() ** exponent)"), + ("await (await dependency())", "await (await dependency())"), + ("await (left + right)", "await (left + right)"), + ], +) +def test_await_expression(source: str, expected: str) -> None: + """Build Await expressions with their correct precedence.""" + node = ast.parse(source, mode="eval").body + expression = get_expression(node, parent=Module("module")) + rendered = str(expression) + + if isinstance(node, ast.Await): + assert isinstance(expression, ExprAwait) + assert rendered == expected + assert ast.dump(ast.parse(rendered, mode="eval").body) == ast.dump(node) + + def test_length_one_tuple_as_string() -> None: """Length-1 tuples must have a trailing comma.""" code = "x = ('a',)"