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

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 

10 

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 

32 

33 

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) 

64 

65 

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 ) 

112 

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) 

122 

123 fixture_teardown_output_value = fixture_teardown_buffer.getvalue() or None 

124 

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 ) 

135 

136 

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) 

151 

152 output_value = teardown_output.getvalue() or None 

153 return session_teardown_failures, output_value 

154 

155 

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 ) 

173 

174 

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 

191 

192 

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 

211 

212 

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 

222 

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) 

234 

235 

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 

262 

263 

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 

274 

275 

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 

319 

320 ( 

321 has_test_failures, 

322 has_fixture_teardown_failures, 

323 has_session_teardown_failures, 

324 ) = has_any_failures(test_results, session_teardown_failures) 

325 

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 

333 

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