Coverage for /home/crpier/Projects/snektest/snektest/fixtures.py: 78%

104 statements  

« prev     ^ index     » next       coverage.py v7.14.1, created at 2026-06-07 00:30 +0300

1import asyncio 

2from collections.abc import AsyncGenerator, Awaitable, Generator, Mapping 

3from dataclasses import dataclass 

4from inspect import isasyncgen, isgenerator 

5from sys import modules 

6from types import CodeType, FunctionType 

7from typing import Any, cast, get_type_hints 

8 

9from snektest.annotations import AsyncSessionFixture, Coroutine, SessionFixture 

10from snektest.models import UnreachableError 

11from snektest.utils import get_code_from_generator, get_func_name_from_generator 

12 

13_SESSION_FIXTURES: dict[ 

14 CodeType, tuple[AsyncGenerator[Any] | Generator[Any] | None, object] 

15] = {} 

16_FUNCTION_FIXTURES: list[AsyncGenerator[Any] | Generator[Any]] = [] 

17 

18 

19def _is_session_fixture_return_annotation(annotation: object) -> bool: 

20 """Return whether an annotation marks a fixture as session-scoped.""" 

21 origin = getattr(annotation, "__origin__", annotation) 

22 return origin in {SessionFixture, AsyncSessionFixture} 

23 

24 

25def _is_session_fixture_function(function: FunctionType) -> bool: 

26 """Return whether a function's return annotation marks a session fixture.""" 

27 try: 

28 return_annotation = get_type_hints(function).get("return") 

29 except Exception: 

30 return False 

31 return return_annotation is not None and _is_session_fixture_return_annotation( 

32 return_annotation 

33 ) 

34 

35 

36def register_session_fixture_from_namespace( 

37 fixture_code: CodeType, 

38 namespace: Mapping[str, object], 

39) -> None: 

40 """Register a matching session fixture function from a namespace.""" 

41 for value in namespace.values(): 

42 if not isinstance(value, FunctionType): 

43 continue 

44 if value.__code__ == fixture_code and _is_session_fixture_function(value): 

45 register_session_fixture(fixture_code) 

46 return 

47 

48 

49def _register_session_fixtures_from_loaded_modules(fixture_code: CodeType) -> None: 

50 """Register matching session fixture functions from already-loaded modules. 

51 

52 Generator objects expose their code object but not the function object that 

53 created them. Searching loaded module globals lets `load_fixture` recover the 

54 fixture function's return annotation without requiring a decorator. 

55 """ 

56 for module in list(modules.values()): 

57 register_session_fixture_from_namespace(fixture_code, vars(module)) 

58 if fixture_code in _SESSION_FIXTURES: 

59 return 

60 

61 

62@dataclass(frozen=True) 

63class _PendingAsyncSessionFixtureSetup: 

64 """Shared async fixture setup while the first load is still pending.""" 

65 

66 awaitable: Awaitable[Any] 

67 

68 

69def _wrap_async_session_fixture_result[R](result: R) -> Coroutine[R]: 

70 async def wrapper() -> R: 

71 return result 

72 

73 return wrapper() 

74 

75 

76def _create_async_session_fixture_setup[R]( 

77 fixture_code: CodeType, 

78 gen: AsyncGenerator[R], 

79) -> Coroutine[R]: 

80 async def result_updater() -> R: 

81 registered_gen, _ = _SESSION_FIXTURES[fixture_code] 

82 if not isasyncgen(registered_gen): 

83 msg = "This should not happen I think" 

84 raise UnreachableError(msg) 

85 result = await anext(registered_gen) 

86 _SESSION_FIXTURES[fixture_code] = (registered_gen, result) 

87 return result 

88 

89 awaitable: Awaitable[R] 

90 try: 

91 loop = asyncio.get_running_loop() 

92 except RuntimeError: 

93 awaitable = result_updater() 

94 else: 

95 awaitable = loop.create_task(result_updater()) 

96 _SESSION_FIXTURES[fixture_code] = ( 

97 gen, 

98 _PendingAsyncSessionFixtureSetup(awaitable), 

99 ) 

100 return cast("Coroutine[R]", awaitable) 

101 

102 

103def register_session_fixture( 

104 fixture_code: CodeType, 

105) -> None: 

106 """Register a session-scoped fixture.""" 

107 if fixture_code not in _SESSION_FIXTURES: 

108 _SESSION_FIXTURES[fixture_code] = (None, None) 

109 

110 

111def get_registered_session_fixtures() -> dict[ 

112 CodeType, tuple[AsyncGenerator[Any] | Generator[Any] | None, object] 

113]: 

114 """Get all registered session fixtures.""" 

115 return _SESSION_FIXTURES 

116 

117 

118def reset_session_fixtures() -> None: 

119 """Clear cached session fixtures for a fresh test run.""" 

120 _SESSION_FIXTURES.clear() 

121 

122 

123def is_session_fixture(fixture_code: CodeType) -> bool: 

124 """Check whether a fixture code object is session-scoped.""" 

125 if fixture_code not in _SESSION_FIXTURES: 

126 _register_session_fixtures_from_loaded_modules(fixture_code) 

127 return fixture_code in _SESSION_FIXTURES 

128 

129 

130def load_session_fixture[R]( 

131 fixture_gen: AsyncGenerator[R] | Generator[R], 

132) -> Coroutine[R] | R: 

133 """Load a session-scoped fixture, creating it on first use and reusing thereafter.""" 

134 fixture_code = get_code_from_generator(fixture_gen) 

135 try: 

136 gen, cached_result = _SESSION_FIXTURES[fixture_code] 

137 except KeyError: 

138 msg = f"Function {fixture_code.co_qualname} was not registered as a session fixture. This shouldn't be possible!" 

139 raise UnreachableError(msg) from None 

140 

141 if gen is None: 

142 gen = fixture_gen 

143 if isasyncgen(gen): 

144 return _create_async_session_fixture_setup(fixture_code, gen) 

145 if isgenerator(gen): 

146 cached_result = next(gen) 

147 _SESSION_FIXTURES[fixture_code] = (gen, cached_result) 

148 return cached_result 

149 msg = "Fixture must be a generator or async generator" 

150 raise UnreachableError(msg) 

151 

152 if isinstance(cached_result, _PendingAsyncSessionFixtureSetup): 

153 return cast("Coroutine[R]", cached_result.awaitable) 

154 if isasyncgen(gen): 

155 return _wrap_async_session_fixture_result(cast("R", cached_result)) 

156 return cast("R", cached_result) 

157 

158 

159def load_function_fixture[R]( 

160 fixture_gen: AsyncGenerator[R] | Generator[R], 

161) -> Coroutine[R] | R: 

162 """Load a function-scoped fixture by registering and yielding its value.""" 

163 if isasyncgen(fixture_gen): 

164 _FUNCTION_FIXTURES.append(fixture_gen) 

165 return anext(fixture_gen) 

166 if isgenerator(fixture_gen): 

167 _FUNCTION_FIXTURES.append(fixture_gen) 

168 return next(fixture_gen) 

169 msg = "Fixture must be a generator or async generator" 

170 raise UnreachableError(msg) 

171 

172 

173def get_active_function_fixtures() -> list[ 

174 tuple[str, AsyncGenerator[Any] | Generator[Any]] 

175]: 

176 """Return the list of active function fixtures, as (function_name, generator) tuples. 

177 

178 Returns: 

179 List of (fixture_name, generator) tuples in reverse order. 

180 """ 

181 fixtures_to_teardown: list[tuple[str, AsyncGenerator[Any] | Generator[Any]]] = [] 

182 # Returning active fixtures in reverse order makes setup/teardown first-in-last-out 

183 for generator in reversed(_FUNCTION_FIXTURES): 

184 fixture_name = get_func_name_from_generator(generator) 

185 fixtures_to_teardown.append((fixture_name, generator)) 

186 

187 _FUNCTION_FIXTURES.clear() 

188 return fixtures_to_teardown