v0.166.0
  1"""
  2Assertion rewriting for test modules.
  3
  4Bare `assert` is the assertion API. When a test module is loaded, each
  5assert in it is rewritten so that a failure knows the values inside the
  6expression: for `assert len(rows) == expected`, what `len(rows)` was, what
  7`rows` was, and what `expected` was.
  8
  9An assert becomes:
 10
 11    a = b = c = NOT_EVALUATED
 12    try:
 13        if not ((a := len((b := rows))) == (c := expected)):
 14            raise failed_assert(..., (a, b, c), message)
 15    finally:
 16        del a, b, c
 17
 18Each part of the expression is wrapped where it stands, in an assignment
 19expression that keeps its value. Nothing moves, so Python evaluates the
 20parts in the order it always would, once each, and leaves out the ones it
 21always would: the right side of an `and` whose left side was false is never
 22evaluated, and is reported as not evaluated.
 23
 24What is raised is an ordinary `AssertionError`, with the message the test
 25gave it or none, as Python would have raised. What was kept travels on it,
 26for the runner to print. It is kept as it was, not printed here: printing a
 27value is the runner's to do, once, with what it knows about the run.
 28
 29The names are deleted whether the assert passes or fails, so an assert
 30holds on to nothing after it. A test that checks an object has been freed
 31would otherwise find it alive, in the hands of the assert before.
 32
 33The expression is shown the way the test file wrote it, taken from the file's
 34own text. Regenerating it from the syntax tree drops parentheses, and
 35`("@" in email) is valid` without them is a different expression.
 36"""
 37
 38import ast
 39import re
 40import textwrap
 41from dataclasses import dataclass
 42from typing import Any
 43
 44__all__ = []
 45
 46# Names put into rewritten test modules. Unique and greppable.
 47_FAILED_ASSERT = "__plain_test_failed_assert__"
 48_NOT_EVALUATED = "__plain_test_not_evaluated__"
 49_KEPT_VALUE = "__plain_test_{index}__"
 50
 51# Where a failed assert's values are kept on the AssertionError.
 52_WATCHED_ASSERT_ATTRIBUTE = "__plain_test_watched_assert__"
 53
 54
 55class _NotEvaluated:
 56    """What a part of an expression holds when Python never evaluated it."""
 57
 58    def __repr__(self) -> str:
 59        return "<not evaluated>"
 60
 61
 62NOT_EVALUATED = _NotEvaluated()
 63
 64
 65@dataclass(frozen=True, kw_only=True)
 66class WatchedValue:
 67    """One part of a failed assert's expression, and what it was."""
 68
 69    # The part as the test file wrote it: `response.status_code`.
 70    source: str
 71    # How far inside the expression it is. The two sides of a comparison
 72    # are 0, what they are made of is 1, and so on.
 73    depth: int
 74    # Whether it is one side of a comparison.
 75    is_operand: bool
 76    # Whether it is written out in the source (`200`, `{"a": 1}`), so that
 77    # printing its value would say the same thing twice.
 78    is_literal: bool
 79    # The value itself, or NOT_EVALUATED.
 80    value: Any
 81
 82
 83@dataclass(frozen=True, kw_only=True)
 84class WatchedAssert:
 85    """A failed assert: its expression and the values inside it."""
 86
 87    # The expression as the test file wrote it, without `assert`.
 88    expression: str
 89    # The parts, outermost first, in the order they are written.
 90    values: tuple[WatchedValue, ...]
 91    # For `assert left == right`, where the two sides are in `values`.
 92    equality: tuple[int, int] | None
 93    # What the test gave after the comma, or None.
 94    message: Any
 95
 96
 97def failed_assert(
 98    expression: str,
 99    parts: tuple[tuple[str, int, bool, bool], ...],
100    equality: tuple[int, int] | None,
101    values: tuple[Any, ...],
102    message: Any,
103) -> AssertionError:
104    """
105    The error a rewritten assert raises. `parts` is what the rewriter knew
106    when the file was compiled, and `values` is what each part was.
107    """
108    error = AssertionError() if message is None else AssertionError(message)
109    watched = WatchedAssert(
110        expression=expression,
111        values=tuple(
112            WatchedValue(
113                source=source,
114                depth=depth,
115                is_operand=is_operand,
116                is_literal=is_literal,
117                value=value,
118            )
119            for (source, depth, is_operand, is_literal), value in zip(
120                parts, values, strict=True
121            )
122        ),
123        equality=equality,
124        message=message,
125    )
126    setattr(error, _WATCHED_ASSERT_ATTRIBUTE, watched)
127    return error
128
129
130def watched_assert_of(error: BaseException) -> WatchedAssert | None:
131    """What a rewritten assert kept, if that is what raised this error."""
132    return getattr(error, _WATCHED_ASSERT_ATTRIBUTE, None)
133
134
135# What ends a line as the parser counts lines. `str.splitlines()` also splits
136# on form feeds and a few other characters, which would put every line after
137# one at the wrong number.
138_LINE_ENDING = re.compile(r"\r\n|\r|\n")
139
140
141def _lines_with_their_endings(source: str) -> list[str]:
142    lines = []
143    start = 0
144    for ending in _LINE_ENDING.finditer(source):
145        lines.append(source[start : ending.end()])
146        start = ending.end()
147    if start < len(source):
148        lines.append(source[start:])
149    return lines
150
151
152def _is_literal(node: ast.expr) -> bool:
153    """Whether an expression is a value written out in full."""
154    if isinstance(node, ast.Constant):
155        return True
156    if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub | ast.UAdd):
157        return isinstance(node.operand, ast.Constant)
158    if isinstance(node, ast.List | ast.Tuple | ast.Set):
159        return all(_is_literal(element) for element in node.elts)
160    if isinstance(node, ast.Dict):
161        return all(
162            key is not None and _is_literal(key) and _is_literal(value)
163            for key, value in zip(node.keys, node.values)
164        )
165    return False
166
167
168# Expressions whose value says nothing a reader can use: a function, and a
169# generator that the call it was passed to has already used up.
170_NEVER_KEPT = (ast.Lambda, ast.GeneratorExp)
171
172# A display's value is what its elements are, which are kept one by one.
173# The whole is kept only where it is one side of a comparison.
174_DISPLAYS = (ast.List, ast.Tuple, ast.Set, ast.Dict)
175
176
177class _Watcher:
178    """
179    Rewrites one assert's expression so that the values inside it are kept,
180    and lists what it kept.
181    """
182
183    def __init__(self, rewriter: _AssertRewriter) -> None:
184        self.rewriter = rewriter
185        # (source, depth, is_operand, is_literal), in the order of `values`.
186        self.parts: list[tuple[str, int, bool, bool]] = []
187
188    def watch(
189        self,
190        node: ast.expr,
191        *,
192        depth: int,
193        is_operand: bool = False,
194        keep: bool = True,
195    ) -> ast.expr:
196        """
197        The expression, rewritten to keep its value and the values inside
198        it. `keep=False` keeps only the ones inside.
199        """
200        if isinstance(node, _NEVER_KEPT):
201            return node
202
203        is_literal = _is_literal(node)
204        if is_literal:
205            # Nothing inside a literal can be anything but what is written.
206            if not is_operand:
207                return node
208            return self._kept(node, node, depth=depth, is_operand=True, is_literal=True)
209
210        if isinstance(node, ast.NamedExpr):
211            # `(total := price())`: what `price()` was is what is wanted,
212            # and the name is the test's own to keep.
213            node.value = self.watch(node.value, depth=depth, is_operand=is_operand)
214            return node
215
216        if isinstance(node, _DISPLAYS) and not is_operand:
217            keep = False
218
219        if not keep:
220            self._watch_inside(node, depth=depth)
221            return node
222
223        # The part is listed before what is inside it, so that the list
224        # reads from the outside in.
225        index = self._list(node, depth=depth, is_operand=is_operand, is_literal=False)
226        self._watch_inside(node, depth=depth + 1)
227        return self._assigned(node, index)
228
229    def _kept(
230        self,
231        original: ast.expr,
232        rewritten: ast.expr,
233        *,
234        depth: int,
235        is_operand: bool,
236        is_literal: bool,
237    ) -> ast.expr:
238        index = self._list(
239            original, depth=depth, is_operand=is_operand, is_literal=is_literal
240        )
241        return self._assigned(rewritten, index)
242
243    def _list(
244        self, node: ast.expr, *, depth: int, is_operand: bool, is_literal: bool
245    ) -> int:
246        self.parts.append(
247            (self.rewriter.source_on_one_line(node), depth, is_operand, is_literal)
248        )
249        return len(self.parts) - 1
250
251    def _assigned(self, node: ast.expr, index: int) -> ast.expr:
252        """`node`, inside an assignment expression that keeps its value."""
253        assigned = ast.NamedExpr(
254            target=ast.Name(id=_KEPT_VALUE.format(index=index), ctx=ast.Store()),
255            value=node,
256        )
257        ast.copy_location(assigned, node)
258        ast.copy_location(assigned.target, node)
259        return assigned
260
261    def _watch_inside(self, node: ast.expr, *, depth: int) -> None:
262        """Rewrite, in place, the expressions that `node` is made of."""
263        match node:
264            case ast.Attribute():
265                node.value = self.watch(node.value, depth=depth)
266            case ast.Subscript():
267                node.value = self.watch(node.value, depth=depth)
268                node.slice = self._watch_slice(node.slice, depth=depth)
269            case ast.Call():
270                self._watch_called(node, depth=depth)
271                node.args = [self._watch_element(arg, depth=depth) for arg in node.args]
272                for keyword in node.keywords:
273                    keyword.value = self.watch(keyword.value, depth=depth)
274            case ast.Compare():
275                node.left = self.watch(node.left, depth=depth, is_operand=True)
276                node.comparators = [
277                    self.watch(comparator, depth=depth, is_operand=True)
278                    for comparator in node.comparators
279                ]
280            case ast.BoolOp():
281                node.values = [self.watch(value, depth=depth) for value in node.values]
282            case ast.UnaryOp():
283                node.operand = self.watch(node.operand, depth=depth)
284            case ast.BinOp():
285                node.left = self.watch(node.left, depth=depth)
286                node.right = self.watch(node.right, depth=depth)
287            case ast.IfExp():
288                # Listed as written, `body if test else orelse`, which is
289                # not the order they are evaluated in.
290                node.body = self.watch(node.body, depth=depth)
291                node.test = self.watch(node.test, depth=depth)
292                node.orelse = self.watch(node.orelse, depth=depth)
293            case ast.Await():
294                # What is awaited is a coroutine, which has nothing to show.
295                # What it was called with does.
296                node.value = self.watch(node.value, depth=depth, keep=False)
297            case ast.List() | ast.Tuple() | ast.Set():
298                node.elts = [
299                    self._watch_element(element, depth=depth) for element in node.elts
300                ]
301            case ast.Dict():
302                node.keys = [
303                    None if key is None else self.watch(key, depth=depth)
304                    for key in node.keys
305                ]
306                node.values = [self.watch(value, depth=depth) for value in node.values]
307            case _:
308                # A name, which has nothing inside it. Or something kept
309                # whole: a comprehension, whose parts are evaluated once for
310                # every item and in a scope of their own; an f-string; and
311                # any expression this doesn't know.
312                pass
313
314    def _watch_called(self, call: ast.Call, *, depth: int) -> None:
315        """
316        Rewrite what is called. The function isn't kept: `len` is `len`. For
317        a method, what it is called on is kept.
318        """
319        called = call.func
320        if isinstance(called, ast.Name):
321            return
322        if isinstance(called, ast.Attribute):
323            called.value = self.watch(called.value, depth=depth)
324            return
325        call.func = self.watch(called, depth=depth)
326
327    def _watch_element(self, node: ast.expr, *, depth: int) -> ast.expr:
328        """An argument, or an element of a display. `*items` keeps `items`."""
329        if isinstance(node, ast.Starred):
330            node.value = self.watch(node.value, depth=depth)
331            return node
332        return self.watch(node, depth=depth)
333
334    def _watch_slice(self, node: ast.expr, *, depth: int) -> ast.expr:
335        """
336        What is between the brackets: `rows[start:stop]` keeps both, and
337        `grid[*position]` keeps `position`.
338        """
339        if isinstance(node, ast.Slice):
340            if node.lower is not None:
341                node.lower = self.watch(node.lower, depth=depth)
342            if node.upper is not None:
343                node.upper = self.watch(node.upper, depth=depth)
344            if node.step is not None:
345                node.step = self.watch(node.step, depth=depth)
346            return node
347        if isinstance(node, ast.Tuple):
348            node.elts = [
349                self._watch_slice(element, depth=depth) for element in node.elts
350            ]
351            return node
352        # Neither a slice nor a starred expression is a value on its own, so
353        # neither can be kept: only what it is made of can.
354        return self._watch_element(node, depth=depth)
355
356
357def _written_out(value: str | int | bool | tuple | None) -> ast.expr:
358    """
359    A value the rewriter knows, as the expression that writes it out. The
360    compiler puts a tuple of constants in with the code's other constants,
361    so an assert that passes builds nothing.
362    """
363    if isinstance(value, tuple):
364        return ast.Tuple(
365            elts=[_written_out(element) for element in value], ctx=ast.Load()
366        )
367    return ast.Constant(value=value)
368
369
370def _is_not(node: ast.expr) -> bool:
371    return isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not)
372
373
374class _AssertRewriter(ast.NodeTransformer):
375    def __init__(self, source: str) -> None:
376        # Split once for the whole file. `ast.get_source_segment()` splits
377        # the source it is given on every call, which for a file with a few
378        # hundred asserts is most of the time collection takes.
379        self.lines = _lines_with_their_endings(source)
380
381    def source_as_written(self, node: ast.expr) -> str:
382        if node.end_lineno is None or node.end_col_offset is None:
383            return ast.unparse(node)
384
385        # Column offsets count UTF-8 bytes, not characters.
386        first_line = self.lines[node.lineno - 1].encode()
387        if node.end_lineno == node.lineno:
388            written = first_line[node.col_offset : node.end_col_offset].decode()
389            return written.strip()
390
391        # An expression written over several lines keeps its shape: the
392        # first line is padded back out to the column it started at, so the
393        # continuation lines keep their indentation relative to it. A tab
394        # stays a tab, since that is what the lines under it are indented by.
395        padding = "".join(
396            character if character == "\t" else " "
397            for character in first_line[: node.col_offset].decode()
398        )
399        last_line = self.lines[node.end_lineno - 1].encode()
400        written = "".join(
401            [
402                padding + first_line[node.col_offset :].decode(),
403                *self.lines[node.lineno : node.end_lineno - 1],
404                last_line[: node.end_col_offset].decode(),
405            ]
406        )
407        return textwrap.dedent(written).strip()
408
409    def source_on_one_line(self, node: ast.expr) -> str:
410        """
411        A part of an expression, to put in front of its value. As written
412        when it was written on one line, and regenerated when it wasn't.
413        """
414        if node.end_lineno == node.lineno:
415            return self.source_as_written(node)
416        return ast.unparse(node)
417
418    def visit_Assert(self, node: ast.Assert) -> list[ast.stmt]:
419        expression = self.source_as_written(node.test)
420        watcher = _Watcher(self)
421
422        # What the whole expression came to is known: it was false. For a
423        # comparison, an `and`, an `or` or a `not`, that says everything
424        # their value could. For anything else (`assert rows`,
425        # `assert response.json_data.get("ok")`) the value is what was
426        # false, and is worth seeing.
427        says_nothing = isinstance(node.test, ast.Compare | ast.BoolOp) or _is_not(
428            node.test
429        )
430        test = watcher.watch(node.test, depth=0, keep=not says_nothing)
431
432        equality = None
433        if (
434            isinstance(node.test, ast.Compare)
435            and len(node.test.ops) == 1
436            and isinstance(node.test.ops[0], ast.Eq)
437        ):
438            # The two sides are the parts at depth 0. A side that is never
439            # kept (a lambda) leaves one, and nothing to compare it with.
440            sides = [
441                index
442                for index, (_, depth, _, _) in enumerate(watcher.parts)
443                if depth == 0
444            ]
445            if len(sides) == 2:
446                equality = (sides[0], sides[1])
447
448        kept_names = [
449            _KEPT_VALUE.format(index=index) for index in range(len(watcher.parts))
450        ]
451        # The test's own expressions go in last. They have their places in
452        # the file already, and the statements made here are given theirs
453        # by walking them, which is a walk these would be most of.
454        in_their_place = ast.Constant(value=None)
455        failure = ast.Call(
456            func=ast.Name(id=_FAILED_ASSERT, ctx=ast.Load()),
457            args=[
458                _written_out(expression),
459                _written_out(tuple(watcher.parts)),
460                _written_out(equality),
461                ast.Tuple(
462                    elts=[ast.Name(id=name, ctx=ast.Load()) for name in kept_names],
463                    ctx=ast.Load(),
464                ),
465                in_their_place,
466            ],
467            keywords=[],
468        )
469        is_false = ast.UnaryOp(op=ast.Not(), operand=in_their_place)
470        check: ast.stmt = ast.If(
471            test=is_false,
472            body=[ast.Raise(exc=failure, cause=None)],
473            orelse=[],
474        )
475        statements: list[ast.stmt]
476        if not kept_names:
477            statements = [check]
478        else:
479            statements = [
480                ast.Assign(
481                    targets=[ast.Name(id=name, ctx=ast.Store()) for name in kept_names],
482                    value=ast.Name(id=_NOT_EVALUATED, ctx=ast.Load()),
483                ),
484                ast.Try(
485                    body=[check],
486                    handlers=[],
487                    orelse=[],
488                    finalbody=[
489                        ast.Delete(
490                            targets=[
491                                ast.Name(id=name, ctx=ast.Del()) for name in kept_names
492                            ]
493                        )
494                    ],
495                ),
496            ]
497
498        for statement in statements:
499            ast.copy_location(statement, node)
500            ast.fix_missing_locations(statement)
501
502        is_false.operand = test
503        if node.msg is not None:
504            # Evaluated where it is, which is only when the assert fails.
505            failure.args[-1] = node.msg
506        return statements
507
508
509def rewrite_asserts(tree: ast.Module, *, source: str) -> ast.Module:
510    """
511    Rewrite the asserts in a parsed test module, and import what they use.
512    `source` is the text `tree` was parsed from.
513    """
514    tree = _AssertRewriter(source).visit(tree)
515
516    # The import goes after any docstring and __future__ imports.
517    insert_at = 0
518    for statement in tree.body:
519        is_docstring = isinstance(statement, ast.Expr) and isinstance(
520            statement.value, ast.Constant
521        )
522        is_future = (
523            isinstance(statement, ast.ImportFrom) and statement.module == "__future__"
524        )
525        if is_docstring or is_future:
526            insert_at += 1
527        else:
528            break
529
530    what_asserts_use = ast.ImportFrom(
531        module="plain.testing.runner.assertions",
532        names=[
533            ast.alias(name="failed_assert", asname=_FAILED_ASSERT),
534            ast.alias(name="NOT_EVALUATED", asname=_NOT_EVALUATED),
535        ],
536        level=0,
537    )
538    # Locate just the injected node — the rewritten asserts already carry
539    # locations, so a whole-tree fix_missing_locations pass isn't needed.
540    if tree.body:
541        ast.copy_location(what_asserts_use, tree.body[0])
542    ast.fix_missing_locations(what_asserts_use)
543    tree.body.insert(insert_at, what_asserts_use)
544    return tree