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