codekingpro/portable-devtools
115k
1import asyncio
2import contextlib
3import inspect
4import warnings
5from typing import (
6 Any,
7 Awaitable,
8 Callable,
9 Dict,
10 Iterator,
11 Optional,
12 Protocol,
13 Union,
14 overload,
15)
16
17import pytest
18
19from .test_utils import (
20 BaseTestServer,
21 RawTestServer,
22 TestClient,
23 TestServer,
24 loop_context,
25 setup_test_loop,
26 teardown_test_loop,
27 unused_port as _unused_port,
28)
29from .web import Application, BaseRequest, Request
30from .web_protocol import _RequestHandler
31
32try:
33 import uvloop
34except ImportError: # pragma: no cover
35 uvloop = None # type: ignore[assignment]
36
37
38class AiohttpClient(Protocol):
39 @overload
40 async def __call__(
41 self,
42 __param: Application,
43 *,
44 server_kwargs: Optional[Dict[str, Any]] = None,
45 **kwargs: Any,
46 ) -> TestClient[Request, Application]: ...
47 @overload
48 async def __call__(
49 self,
50 __param: BaseTestServer,
51 *,
52 server_kwargs: Optional[Dict[str, Any]] = None,
53 **kwargs: Any,
54 ) -> TestClient[BaseRequest, None]: ...
55
56
57class AiohttpServer(Protocol):
58 def __call__(
59 self, app: Application, *, port: Optional[int] = None, **kwargs: Any
60 ) -> Awaitable[TestServer]: ...
61
62
63class AiohttpRawServer(Protocol):
64 def __call__(
65 self, handler: _RequestHandler, *, port: Optional[int] = None, **kwargs: Any
66 ) -> Awaitable[RawTestServer]: ...
67
68
69def pytest_addoption(parser): # type: ignore[no-untyped-def]
70 parser.addoption(
71 "--aiohttp-fast",
72 action="store_true",
73 default=False,
74 help="run tests faster by disabling extra checks",
75 )
76 parser.addoption(
77 "--aiohttp-loop",
78 action="store",
79 default="pyloop",
80 help="run tests with specific loop: pyloop, uvloop or all",
81 )
82 parser.addoption(
83 "--aiohttp-enable-loop-debug",
84 action="store_true",
85 default=False,
86 help="enable event loop debug mode",
87 )
88
89
90def pytest_fixture_setup(fixturedef): # type: ignore[no-untyped-def]
91 """Set up pytest fixture.
92
93 Allow fixtures to be coroutines. Run coroutine fixtures in an event loop.
94 """
95 func = fixturedef.func
96
97 if inspect.isasyncgenfunction(func):
98 # async generator fixture
99 is_async_gen = True
100 elif inspect.iscoroutinefunction(func):
101 # regular async fixture
102 is_async_gen = False
103 else:
104 # not an async fixture, nothing to do
105 return
106
107 strip_request = False
108 if "request" not in fixturedef.argnames:
109 fixturedef.argnames += ("request",)
110 strip_request = True
111
112 def wrapper(*args, **kwargs): # type: ignore[no-untyped-def]
113 request = kwargs["request"]
114 if strip_request:
115 del kwargs["request"]
116
117 # if neither the fixture nor the test use the 'loop' fixture,
118 # 'getfixturevalue' will fail because the test is not parameterized
119 # (this can be removed someday if 'loop' is no longer parameterized)
120 if "loop" not in request.fixturenames:
121 raise Exception(
122 "Asynchronous fixtures must depend on the 'loop' fixture or "
123 "be used in tests depending from it."
124 )
125
126 _loop = request.getfixturevalue("loop")
127
128 if is_async_gen:
129 # for async generators, we need to advance the generator once,
130 # then advance it again in a finalizer
131 gen = func(*args, **kwargs)
132
133 def finalizer(): # type: ignore[no-untyped-def]
134 try:
135 return _loop.run_until_complete(gen.__anext__())
136 except StopAsyncIteration:
137 pass
138
139 request.addfinalizer(finalizer)
140 return _loop.run_until_complete(gen.__anext__())
141 else:
142 return _loop.run_until_complete(func(*args, **kwargs))
143
144 fixturedef.func = wrapper
145
146
147@pytest.fixture
148def fast(request): # type: ignore[no-untyped-def]
149 """--fast config option"""
150 return request.config.getoption("--aiohttp-fast")
151
152
153@pytest.fixture
154def loop_debug(request): # type: ignore[no-untyped-def]
155 """--enable-loop-debug config option"""
156 return request.config.getoption("--aiohttp-enable-loop-debug")
157
158
159@contextlib.contextmanager
160def _runtime_warning_context(): # type: ignore[no-untyped-def]
161 """Context manager which checks for RuntimeWarnings.
162
163 This exists specifically to
164 avoid "coroutine 'X' was never awaited" warnings being missed.
165
166 If RuntimeWarnings occur in the context a RuntimeError is raised.
167 """
168 with warnings.catch_warnings(record=True) as _warnings:
169 yield
170 rw = [
171 "{w.filename}:{w.lineno}:{w.message}".format(w=w)
172 for w in _warnings
173 if w.category == RuntimeWarning
174 ]
175 if rw:
176 raise RuntimeError(
177 "{} Runtime Warning{},\n{}".format(
178 len(rw), "" if len(rw) == 1 else "s", "\n".join(rw)
179 )
180 )
181
182
183@contextlib.contextmanager
184def _passthrough_loop_context(loop, fast=False): # type: ignore[no-untyped-def]
185 """Passthrough loop context.
186
187 Sets up and tears down a loop unless one is passed in via the loop
188 argument when it's passed straight through.
189 """
190 if loop:
191 # loop already exists, pass it straight through
192 yield loop
193 else:
194 # this shadows loop_context's standard behavior
195 loop = setup_test_loop()
196 yield loop
197 teardown_test_loop(loop, fast=fast)
198
199
200def pytest_pycollect_makeitem(collector, name, obj): # type: ignore[no-untyped-def]
201 """Fix pytest collecting for coroutines."""
202 if collector.funcnamefilter(name) and inspect.iscoroutinefunction(obj):
203 return list(collector._genfunctions(name, obj))
204
205
206def pytest_pyfunc_call(pyfuncitem): # type: ignore[no-untyped-def]
207 """Run coroutines in an event loop instead of a normal function call."""
208 fast = pyfuncitem.config.getoption("--aiohttp-fast")
209 if inspect.iscoroutinefunction(pyfuncitem.function):
210 existing_loop = (
211 pyfuncitem.funcargs.get("proactor_loop")
212 or pyfuncitem.funcargs.get("selector_loop")
213 or pyfuncitem.funcargs.get("uvloop_loop")
214 or pyfuncitem.funcargs.get("loop", None)
215 )
216
217 with _runtime_warning_context():
218 with _passthrough_loop_context(existing_loop, fast=fast) as _loop:
219 testargs = {
220 arg: pyfuncitem.funcargs[arg]
221 for arg in pyfuncitem._fixtureinfo.argnames
222 }
223 _loop.run_until_complete(pyfuncitem.obj(**testargs))
224
225 return True
226
227
228def pytest_generate_tests(metafunc): # type: ignore[no-untyped-def]
229 if "loop_factory" not in metafunc.fixturenames:
230 return
231
232 loops = metafunc.config.option.aiohttp_loop
233 avail_factories: dict[str, Callable[[], asyncio.AbstractEventLoop]]
234 avail_factories = {"pyloop": asyncio.new_event_loop}
235
236 if uvloop is not None: # pragma: no cover
237 avail_factories["uvloop"] = uvloop.new_event_loop
238
239 if loops == "all":
240 loops = "pyloop,uvloop?"
241
242 factories = {} # type: ignore[var-annotated]
243 for name in loops.split(","):
244 required = not name.endswith("?")
245 name = name.strip(" ?")
246 if name not in avail_factories: # pragma: no cover
247 if required:
248 raise ValueError(
249 "Unknown loop '%s', available loops: %s"
250 % (name, list(factories.keys()))
251 )
252 else:
253 continue
254 factories[name] = avail_factories[name]
255 metafunc.parametrize(
256 "loop_factory", list(factories.values()), ids=list(factories.keys())
257 )
258
259
260@pytest.fixture
261def loop(
262 loop_factory: Callable[[], asyncio.AbstractEventLoop],
263 fast: bool,
264 loop_debug: bool,
265) -> Iterator[asyncio.AbstractEventLoop]:
266 """Return an instance of the event loop."""
267 with loop_context(loop_factory, fast=fast) as _loop:
268 if loop_debug:
269 _loop.set_debug(True) # pragma: no cover
270 asyncio.set_event_loop(_loop)
271 yield _loop
272
273
274@pytest.fixture
275def proactor_loop() -> Iterator[asyncio.AbstractEventLoop]:
276 factory = asyncio.ProactorEventLoop # type: ignore[attr-defined]
277
278 with loop_context(factory) as _loop:
279 asyncio.set_event_loop(_loop)
280 yield _loop
281
282
283@pytest.fixture
284def unused_port(aiohttp_unused_port: Callable[[], int]) -> Callable[[], int]:
285 warnings.warn(
286 "Deprecated, use aiohttp_unused_port fixture instead",
287 DeprecationWarning,
288 stacklevel=2,
289 )
290 return aiohttp_unused_port
291
292
293@pytest.fixture
294def aiohttp_unused_port() -> Callable[[], int]:
295 """Return a port that is unused on the current host."""
296 return _unused_port
297
298
299@pytest.fixture
300def aiohttp_server(loop: asyncio.AbstractEventLoop) -> Iterator[AiohttpServer]:
301 """Factory to create a TestServer instance, given an app.
302
303 aiohttp_server(app, **kwargs)
304 """
305 servers = []
306
307 async def go(
308 app: Application,
309 *,
310 host: str = "127.0.0.1",
311 port: Optional[int] = None,
312 **kwargs: Any,
313 ) -> TestServer:
314 server = TestServer(app, host=host, port=port)
315 await server.start_server(loop=loop, **kwargs)
316 servers.append(server)
317 return server
318
319 yield go
320
321 async def finalize() -> None:
322 while servers:
323 await servers.pop().close()
324
325 loop.run_until_complete(finalize())
326
327
328@pytest.fixture
329def test_server(aiohttp_server): # type: ignore[no-untyped-def] # pragma: no cover
330 warnings.warn(
331 "Deprecated, use aiohttp_server fixture instead",
332 DeprecationWarning,
333 stacklevel=2,
334 )
335 return aiohttp_server
336
337
338@pytest.fixture
339def aiohttp_raw_server(loop: asyncio.AbstractEventLoop) -> Iterator[AiohttpRawServer]:
340 """Factory to create a RawTestServer instance, given a web handler.
341
342 aiohttp_raw_server(handler, **kwargs)
343 """
344 servers = []
345
346 async def go(
347 handler: _RequestHandler, *, port: Optional[int] = None, **kwargs: Any
348 ) -> RawTestServer:
349 server = RawTestServer(handler, port=port)
350 await server.start_server(loop=loop, **kwargs)
351 servers.append(server)
352 return server
353
354 yield go
355
356 async def finalize() -> None:
357 while servers:
358 await servers.pop().close()
359
360 loop.run_until_complete(finalize())
361
362
363@pytest.fixture
364def raw_test_server( # type: ignore[no-untyped-def] # pragma: no cover
365 aiohttp_raw_server,
366):
367 warnings.warn(
368 "Deprecated, use aiohttp_raw_server fixture instead",
369 DeprecationWarning,
370 stacklevel=2,
371 )
372 return aiohttp_raw_server
373
374
375@pytest.fixture
376def aiohttp_client(loop: asyncio.AbstractEventLoop) -> Iterator[AiohttpClient]:
377 """Factory to create a TestClient instance.
378
379 aiohttp_client(app, **kwargs)
380 aiohttp_client(server, **kwargs)
381 aiohttp_client(raw_server, **kwargs)
382 """
383 clients = []
384
385 @overload
386 async def go(
387 __param: Application,
388 *,
389 server_kwargs: Optional[Dict[str, Any]] = None,
390 **kwargs: Any,
391 ) -> TestClient[Request, Application]: ...
392
393 @overload
394 async def go(
395 __param: BaseTestServer,
396 *,
397 server_kwargs: Optional[Dict[str, Any]] = None,
398 **kwargs: Any,
399 ) -> TestClient[BaseRequest, None]: ...
400
401 async def go(
402 __param: Union[Application, BaseTestServer],
403 *args: Any,
404 server_kwargs: Optional[Dict[str, Any]] = None,
405 **kwargs: Any,
406 ) -> TestClient[Any, Any]:
407 if isinstance(__param, Callable) and not isinstance( # type: ignore[arg-type]
408 __param, (Application, BaseTestServer)
409 ):
410 __param = __param(loop, *args, **kwargs)
411 kwargs = {}
412 else:
413 assert not args, "args should be empty"
414
415 if isinstance(__param, Application):
416 server_kwargs = server_kwargs or {}
417 server = TestServer(__param, loop=loop, **server_kwargs)
418 client = TestClient(server, loop=loop, **kwargs)
419 elif isinstance(__param, BaseTestServer):
420 client = TestClient(__param, loop=loop, **kwargs)
421 else:
422 raise ValueError("Unknown argument type: %r" % type(__param))
423
424 await client.start_server()
425 clients.append(client)
426 return client
427
428 yield go
429
430 async def finalize() -> None:
431 while clients:
432 await clients.pop().close()
433
434 loop.run_until_complete(finalize())
435
436
437@pytest.fixture
438def test_client(aiohttp_client): # type: ignore[no-untyped-def] # pragma: no cover
439 warnings.warn(
440 "Deprecated, use aiohttp_client fixture instead",
441 DeprecationWarning,
442 stacklevel=2,
443 )
444 return aiohttp_client
445 