Coverage for /home/crpier/Projects/snektest/snektest/execution.py: 52%
169 statements
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 00:30 +0300
« prev ^ index » next coverage.py v7.14.1, created at 2026-06-07 00:30 +0300
1import asyncio
2import pdb # noqa: T100
3import sys
4import time
5from collections.abc import Callable
6from inspect import isasyncgen, iscoroutine, isgenerator
7from pathlib import Path
8from types import TracebackType
9from typing import cast
11from snektest.annotations import Coroutine
12from snektest.collection import TestsQueue
13from snektest.fixtures import (
14 get_active_function_fixtures,
15 get_registered_session_fixtures,
16 reset_session_fixtures,
17)
18from snektest.models import (
19 AssertionFailure,
20 BadRequestError,
21 ErrorResult,
22 FailedResult,
23 PassedResult,
24 TeardownFailure,
25 TestName,
26 TestResult,
27 UnreachableError,
28)
29from snektest.output import maybe_capture_output
30from snektest.presenter import print_failures, print_summary, print_test_result
31from snektest.utils import get_test_function_markers, get_test_function_params
34async def teardown_fixture(
35 fixture_name: str,
36 generator: object,
37 *,
38 exc_info_provider: Callable[
39 [], tuple[object | None, object | None, TracebackType | None]
40 ] = sys.exc_info,
41) -> TeardownFailure | None:
42 """Teardown a single fixture and return failure if it occurs."""
43 try:
44 if isasyncgen(generator):
45 await anext(generator)
46 elif isgenerator(generator):
47 next(generator)
48 except StopAsyncIteration, StopIteration:
49 return None
50 except Exception:
51 exc_type, exc_value, traceback = exc_info_provider()
52 if exc_type is None or exc_value is None or traceback is None:
53 msg = "Invalid exception info gathered during teardown. This shouldn't be possible!"
54 raise UnreachableError(msg) from None
55 return TeardownFailure(
56 fixture_name=fixture_name,
57 exc_type=cast("type[BaseException]", exc_type),
58 exc_value=cast("BaseException", exc_value),
59 traceback=traceback,
60 )
61 else:
62 msg = f"Incorrect fixture function {fixture_name} yielded more than once"
63 raise BadRequestError(msg)
66async def execute_test(
67 name: TestName,
68 func: Callable[..., Coroutine[None] | None],
69 *,
70 capture_output: bool = True,
71 exc_info_provider: Callable[
72 [], tuple[object | None, object | None, TracebackType | None]
73 ] = sys.exc_info,
74) -> TestResult:
75 """Execute a single test function with fixtures and output capture."""
76 with maybe_capture_output(capture_output) as (output_buffer, captured_warnings):
77 param_values = ()
78 if name.params_part:
79 param_values = [
80 param.value
81 for param in get_test_function_params(func)[name.params_part]
82 ]
83 test_start = time.monotonic()
84 try:
85 res = func(*param_values)
86 if iscoroutine(res):
87 await res
88 duration = time.monotonic() - test_start
89 result = PassedResult()
90 except (AssertionFailure, asyncio.CancelledError):
91 duration = time.monotonic() - test_start
92 exc_type, exc_value, traceback = exc_info_provider()
93 if exc_type is None or exc_value is None or traceback is None:
94 msg = "Invalid exception info gathered. This shouldn't be possible!"
95 raise UnreachableError(msg) from None
96 result = FailedResult(
97 exc_type=cast("type[BaseException]", exc_type),
98 exc_value=cast("BaseException", exc_value),
99 traceback=traceback,
100 )
101 except Exception:
102 duration = time.monotonic() - test_start
103 exc_type, exc_value, traceback = exc_info_provider()
104 if exc_type is None or exc_value is None or traceback is None:
105 msg = "Invalid exception info gathered. This shouldn't be possible!"
106 raise UnreachableError(msg) from None
107 result = ErrorResult(
108 exc_type=cast("type[BaseException]", exc_type),
109 exc_value=cast("BaseException", exc_value),
110 traceback=traceback,
111 )
113 with maybe_capture_output(capture_output) as (
114 fixture_teardown_buffer,
115 _,
116 ):
117 fixture_teardown_failures: list[TeardownFailure] = []
118 for fixture_name, generator in get_active_function_fixtures():
119 failure = await teardown_fixture(fixture_name, generator)
120 if failure:
121 fixture_teardown_failures.append(failure)
123 fixture_teardown_output_value = fixture_teardown_buffer.getvalue() or None
125 return TestResult(
126 name=name,
127 duration=duration,
128 result=result,
129 markers=get_test_function_markers(func),
130 captured_output=output_buffer,
131 fixture_teardown_failures=fixture_teardown_failures,
132 fixture_teardown_output=fixture_teardown_output_value,
133 warnings=captured_warnings,
134 )
137async def teardown_session_fixtures(
138 *, capture_output: bool
139) -> tuple[list[TeardownFailure], str | None]:
140 """Teardown all session fixtures and return failures and output."""
141 with maybe_capture_output(capture_output) as (teardown_output, _):
142 session_teardown_failures: list[TeardownFailure] = []
143 for fixture_func, (gen, _) in reversed(
144 get_registered_session_fixtures().items()
145 ):
146 if gen is not None:
147 fixture_name = fixture_func.co_name
148 failure = await teardown_fixture(fixture_name, gen)
149 if failure:
150 session_teardown_failures.append(failure)
152 output_value = teardown_output.getvalue() or None
153 return session_teardown_failures, output_value
156def has_any_failures(
157 test_results: list[TestResult], session_teardown_failures: list[TeardownFailure]
158) -> tuple[bool, bool, bool]:
159 """Check for test failures, fixture failures, and session failures."""
160 has_test_failures = any(
161 isinstance(result.result, (FailedResult, ErrorResult))
162 for result in test_results
163 )
164 has_fixture_teardown_failures = any(
165 result.fixture_teardown_failures for result in test_results
166 )
167 has_session_teardown_failures = len(session_teardown_failures) > 0
168 return (
169 has_test_failures,
170 has_fixture_teardown_failures,
171 has_session_teardown_failures,
172 )
175def _resolve_path(
176 path: Path | None,
177 *,
178 resolver: Callable[[Path], Path] = Path.resolve,
179) -> Path | None:
180 if path is None:
181 return None
182 try:
183 resolved = resolver(path)
184 except FileNotFoundError:
185 return path
186 if resolved is path:
187 return resolved
188 if str(resolved):
189 return resolved
190 return resolved
193def _trim_traceback(
194 traceback: TracebackType, *, stop_at: TracebackType
195) -> TracebackType:
196 frames: list[TracebackType] = []
197 current = traceback
198 while current is not None:
199 frames.append(current)
200 if current is stop_at:
201 break
202 current = current.tb_next
203 new_traceback: TracebackType | None = None
204 for frame in reversed(frames):
205 new_traceback = TracebackType(
206 new_traceback, frame.tb_frame, frame.tb_lasti, frame.tb_lineno
207 )
208 if new_traceback is None:
209 return traceback
210 return new_traceback
213def _traceback_for_file(
214 traceback: TracebackType,
215 *,
216 preferred_file: Path | None,
217 resolver: Callable[[Path], Path] = Path.resolve,
218) -> TracebackType:
219 preferred = _resolve_path(preferred_file, resolver=resolver)
220 if preferred is None:
221 return traceback
223 selected: TracebackType | None = None
224 current = traceback
225 while current is not None:
226 frame_path = Path(current.tb_frame.f_code.co_filename)
227 resolved = _resolve_path(frame_path, resolver=resolver)
228 if resolved == preferred:
229 selected = current
230 current = current.tb_next
231 if selected is None:
232 return traceback
233 return _trim_traceback(traceback, stop_at=selected)
236def _maybe_debug_test_result(
237 test_result: TestResult,
238 *,
239 pdb_on_failure: bool,
240 post_mortem: Callable[[TracebackType], None] = pdb.post_mortem,
241 resolver: Callable[[Path], Path] = Path.resolve,
242) -> bool:
243 if not pdb_on_failure:
244 return False
245 if isinstance(test_result.result, (FailedResult, ErrorResult)):
246 traceback = _traceback_for_file(
247 test_result.result.traceback,
248 preferred_file=test_result.name.file_path,
249 resolver=resolver,
250 )
251 post_mortem(traceback)
252 return True
253 if test_result.fixture_teardown_failures:
254 traceback = _traceback_for_file(
255 test_result.fixture_teardown_failures[0].traceback,
256 preferred_file=test_result.name.file_path,
257 resolver=resolver,
258 )
259 post_mortem(traceback)
260 return True
261 return False
264def _maybe_debug_session_teardown(
265 session_teardown_failures: list[TeardownFailure],
266 *,
267 pdb_on_failure: bool,
268 post_mortem: Callable[[TracebackType], None] = pdb.post_mortem,
269) -> bool:
270 if not pdb_on_failure or not session_teardown_failures:
271 return False
272 post_mortem(session_teardown_failures[0].traceback)
273 return True
276async def run_tests( # noqa: PLR0913
277 queue: TestsQueue,
278 *,
279 capture_output: bool = True,
280 pdb_on_failure: bool = False,
281 collection_failed: Callable[[], bool] = lambda: False,
282 post_mortem: Callable[[TracebackType], None] = pdb.post_mortem,
283 resolver: Callable[[Path], Path] = Path.resolve,
284) -> tuple[list[TestResult], list[TeardownFailure]]:
285 """Run all tests from the queue and handle session fixture teardown."""
286 total_duration = time.monotonic()
287 test_results: list[TestResult] = []
288 session_teardown_failures: list[TeardownFailure] = []
289 pdb_triggered = False
290 try:
291 while True:
292 name, func = await queue.get()
293 test_result = await execute_test(name, func, capture_output=capture_output)
294 test_results.append(test_result)
295 print_test_result(test_result)
296 if not pdb_triggered and _maybe_debug_test_result(
297 test_result,
298 pdb_on_failure=pdb_on_failure,
299 post_mortem=post_mortem,
300 resolver=resolver,
301 ):
302 pdb_triggered = True
303 break
304 except asyncio.QueueShutDown:
305 pass
306 finally:
307 if collection_failed():
308 reset_session_fixtures()
309 else:
310 session_teardown_failures, session_output = await teardown_session_fixtures(
311 capture_output=capture_output
312 )
313 if not pdb_triggered and _maybe_debug_session_teardown(
314 session_teardown_failures,
315 pdb_on_failure=pdb_on_failure,
316 post_mortem=post_mortem,
317 ):
318 pdb_triggered = True
320 (
321 has_test_failures,
322 has_fixture_teardown_failures,
323 has_session_teardown_failures,
324 ) = has_any_failures(test_results, session_teardown_failures)
326 session_output_for_display = None
327 if session_output and (
328 has_test_failures
329 or has_fixture_teardown_failures
330 or has_session_teardown_failures
331 ):
332 session_output_for_display = session_output
334 print_failures(
335 test_results,
336 session_teardown_failures=session_teardown_failures,
337 session_teardown_output=session_output_for_display,
338 )
339 print_summary(
340 test_results,
341 session_teardown_failures=session_teardown_failures,
342 total_duration=time.monotonic() - total_duration,
343 )
344 reset_session_fixtures()
345 return test_results, session_teardown_failures