1"""
2Declarative test decorators.
3
4Decorators declare static facts about a test — they never inject runtime
5values or alter control flow. The test runner (plain.testing.runner) reads the
6attributes they attach.
7"""
8
9import enum
10from collections.abc import Callable
11from typing import Any
12
13from .definition import TestDefinitionError
14
15__all__ = ["case", "cases", "skip", "tag"]
16
17# Attribute names the runner reads. Unique and greppable on purpose.
18TEST_CASES_ATTRIBUTE = "__plain_test_cases__"
19TEST_SKIP_ATTRIBUTE = "__plain_test_skip__"
20TEST_TAGS_ATTRIBUTE = "__plain_test_tags__"
21
22
23class case:
24 """
25 One `@cases` entry with a name of its own.
26
27 @cases(
28 case("[email protected]", True, id="plain address"),
29 case("nope", False, id="no at sign"),
30 )
31 def test_email_validation(email, valid):
32 assert is_valid_email(email) is valid
33
34 The id becomes the test id — `test_email_validation[no at sign]` — so a
35 failure and its re-run command name the case for what it is about. The
36 id sits on the case it names, so adding or reordering cases can't
37 quietly shift the names onto the wrong values.
38 """
39
40 __slots__ = ("id", "values")
41
42 def __init__(self, *values: Any, id: str) -> None:
43 if not id or not id.strip():
44 raise TestDefinitionError("case() requires a non-empty id")
45 self.values = values
46 self.id = id
47
48 def __repr__(self) -> str:
49 arguments = ", ".join(repr(value) for value in self.values)
50 return f"case({arguments}, id={self.id!r})"
51
52
53# A case's id is put between brackets after the test's name, and typed after
54# `plain test`. Past this many characters it says less than a number would.
55_LONGEST_ID_FROM_VALUES = 60
56
57
58def cases(*case_args: Any) -> Callable:
59 """
60 Run a test once for each case. Each argument is a case, passed as the
61 test function's positional arguments.
62
63 @cases(
64 ("[email protected]", True),
65 ("nope", False),
66 )
67 def test_email_validation(email, valid):
68 assert is_valid_email(email) is valid
69
70 A non-tuple case is passed as a single argument.
71
72 A case is reported by its values, joined with `-`:
73 `test_email_validation[nope-False]`. That needs every value of every
74 case to be a string, a number, a boolean, `None` or an enum member, and
75 every case to come out different. Otherwise the cases are numbered,
76 `test_email_validation[0]`. Wrap a case in `case(..., id="...")` to give
77 it a name of its own.
78
79 A test takes one `@cases`, and each case is one flat tuple: the test's
80 values in the order of its parameters. For every combination of two
81 lists, build the cases from both:
82
83 @cases(*[(a, b, c) for a in FIRST for b, c in SECOND])
84
85 Each `for` names what one entry of its list holds.
86 """
87 entries: list[tuple[tuple[Any, ...], str | None]] = []
88 for entry in case_args:
89 if isinstance(entry, case):
90 entries.append((entry.values, entry.id))
91 elif isinstance(entry, tuple):
92 entries.append((entry, None))
93 else:
94 entries.append(((entry,), None))
95
96 if not entries:
97 raise TestDefinitionError("cases() requires at least one case")
98
99 given_ids = [given_id for _, given_id in entries if given_id is not None]
100 duplicates = {given_id for given_id in given_ids if given_ids.count(given_id) > 1}
101 if duplicates:
102 raise TestDefinitionError(
103 f"cases() ids must be unique — repeated: {sorted(duplicates)}"
104 )
105
106 ids = _ids_from_values(entries)
107 if ids is None:
108 ids = _ids_from_positions(entries)
109 normalized = [(values, id) for (values, _), id in zip(entries, ids, strict=True)]
110
111 def decorator(func: Callable) -> Callable:
112 if TEST_CASES_ATTRIBUTE in vars(func):
113 name = getattr(func, "__name__", "this test")
114 raise TestDefinitionError(
115 f"{name} already has @cases. A second one would "
116 "replace the first, not combine with it. Write the "
117 "combinations as one @cases. Each case is one flat tuple: "
118 "the test's values, in the order of its parameters.\n"
119 "\n"
120 " @cases(*[(a, b, c) for a in FIRST for b, c in SECOND])\n"
121 f" def {name}(a, b, c): ...\n"
122 "\n"
123 "Each `for` names what one entry of its list holds: "
124 "`for a in FIRST` when the entries are single values, "
125 "`for b, c in SECOND` when they are tuples. "
126 "(`itertools.product` would nest those: `(a, (b, c))`.)"
127 )
128 setattr(func, TEST_CASES_ATTRIBUTE, normalized)
129 return func
130
131 return decorator
132
133
134def _ids_from_values(
135 entries: list[tuple[tuple[Any, ...], str | None]],
136) -> list[str] | None:
137 """
138 An id for every case: the one it was given, or its values joined with
139 `-`. None when a case's values can't be written that way, or when two
140 cases would come out the same. Then none of them are named for their
141 values, so that one test's cases are all named or all numbered.
142 """
143 ids = []
144 for values, given_id in entries:
145 if given_id is not None:
146 ids.append(given_id)
147 continue
148 written = [_written_in_an_id(value) for value in values]
149 if not written or None in written:
150 return None
151 from_values = "-".join(part for part in written if part is not None)
152 if len(from_values) > _LONGEST_ID_FROM_VALUES:
153 return None
154 ids.append(from_values)
155
156 if len(set(ids)) != len(ids):
157 return None
158 return ids
159
160
161def _written_in_an_id(value: Any) -> str | None:
162 """A value as it is written in a case's id, or None if it can't be."""
163 if isinstance(value, enum.Enum):
164 return value.name
165 if value is None or type(value) in (bool, int, float):
166 return repr(value)
167 # A string has to be readable where it is printed, and the same when it
168 # is typed back: nothing that isn't a character to look at, and no space
169 # at either end to be lost.
170 if type(value) is str and value and value.isprintable() and value == value.strip():
171 return value
172 return None
173
174
175def _ids_from_positions(
176 entries: list[tuple[tuple[Any, ...], str | None]],
177) -> list[str]:
178 """An id for every case: the one it was given, or its number."""
179 ids = []
180 for position, (_, given_id) in enumerate(entries):
181 ids.append(given_id if given_id is not None else str(position))
182
183 for position, (_, given_id) in enumerate(entries):
184 if given_id is None:
185 continue
186 if ids.count(given_id) > 1:
187 numbered = ids.index(given_id)
188 if numbered == position:
189 numbered = ids.index(given_id, position + 1)
190 raise TestDefinitionError(
191 f"cases() ids must be unique — case {position} is named "
192 f"{given_id!r}, which is what case {numbered} is numbered. "
193 "Give it another name."
194 )
195 return ids
196
197
198def skip(reason: str) -> Callable:
199 """
200 Always skip this test, with the reason shown in the report.
201
202 To skip from inside a running test, call `skip_test(reason)`.
203 """
204 # A bare `@skip` would hand the test function in as the reason, and the
205 # test would silently stop existing.
206 if not isinstance(reason, str) or not reason.strip():
207 raise TestDefinitionError('@skip requires a reason: @skip("why")')
208
209 def decorator(func: Callable) -> Callable:
210 setattr(func, TEST_SKIP_ATTRIBUTE, reason)
211 return func
212
213 return decorator
214
215
216def tag(*names: str) -> Callable:
217 """
218 Label a test for selection (`plain test --tag slow`) or for package
219 lifecycles that change behavior per-test (e.g. `@isolated_db` from
220 plain.postgres is a tag under the hood).
221 """
222
223 # A bare `@tag` would hand the test function in as a name, and the test
224 # would silently stop existing.
225 if not names or not all(isinstance(name, str) and name.strip() for name in names):
226 raise TestDefinitionError('@tag requires at least one name: @tag("slow")')
227
228 def decorator(func: Callable) -> Callable:
229 existing = getattr(func, TEST_TAGS_ATTRIBUTE, ())
230 setattr(func, TEST_TAGS_ATTRIBUTE, (*existing, *names))
231 return func
232
233 return decorator