codekingpro/portable-devtools
115k
1"""Utilities shared by tests."""
2
3import asyncio
4import contextlib
5import gc
6import inspect
7import ipaddress
8import os
9import socket
10import sys
11import warnings
12from abc import ABC, abstractmethod
13from types import TracebackType
14from typing import (
15 TYPE_CHECKING,
16 Any,
17 Callable,
18 Generic,
19 Iterator,
20 List,
21 Optional,
22 Type,
23 TypeVar,
24 cast,
25 overload,
26)
27from unittest import IsolatedAsyncioTestCase, mock
28
29from aiosignal import Signal
30from multidict import CIMultiDict, CIMultiDictProxy
31from yarl import URL
32
33import aiohttp
34from aiohttp.client import (
35 _RequestContextManager,
36 _RequestOptions,
37 _WSRequestContextManager,
38)
39
40from . import ClientSession, hdrs
41from .abc import AbstractCookieJar
42from .client_reqrep import ClientResponse
43from .client_ws import ClientWebSocketResponse
44from .helpers import sentinel
45from .http import HttpVersion, RawRequestMessage
46from .streams import EMPTY_PAYLOAD, StreamReader
47from .typedefs import StrOrURL
48from .web import (
49 Application,
50 AppRunner,
51 BaseRequest,
52 BaseRunner,
53 Request,
54 Server,
55 ServerRunner,
56 SockSite,
57 UrlMappingMatchInfo,
58)
59from .web_protocol import _RequestHandler
60
61if TYPE_CHECKING:
62 from ssl import SSLContext
63else:
64 SSLContext = None
65
66if sys.version_info >= (3, 11) and TYPE_CHECKING:
67 from typing import Unpack
68
69if sys.version_info >= (3, 11):
70 from typing import Self
71else:
72 Self = Any
73
74_ApplicationNone = TypeVar("_ApplicationNone", Application, None)
75_Request = TypeVar("_Request", bound=BaseRequest)
76
77REUSE_ADDRESS = os.name == "posix" and sys.platform != "cygwin"
78
79
80def get_unused_port_socket(
81 host: str, family: socket.AddressFamily = socket.AF_INET
82) -> socket.socket:
83 return get_port_socket(host, 0, family)
84
85
86def get_port_socket(
87 host: str, port: int, family: socket.AddressFamily
88) -> socket.socket:
89 s = socket.socket(family, socket.SOCK_STREAM)
90 if REUSE_ADDRESS:
91 # Windows has different semantics for SO_REUSEADDR,
92 # so don't set it. Ref:
93 # https://docs.microsoft.com/en-us/windows/win32/winsock/using-so-reuseaddr-and-so-exclusiveaddruse
94 s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
95 s.bind((host, port))
96 return s
97
98
99def unused_port() -> int:
100 """Return a port that is unused on the current host."""
101 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
102 s.bind(("127.0.0.1", 0))
103 return cast(int, s.getsockname()[1])
104
105
106class BaseTestServer(ABC):
107 __test__ = False
108
109 def __init__(
110 self,
111 *,
112 scheme: str = "",
113 loop: Optional[asyncio.AbstractEventLoop] = None,
114 host: str = "127.0.0.1",
115 port: Optional[int] = None,
116 skip_url_asserts: bool = False,
117 socket_factory: Callable[
118 [str, int, socket.AddressFamily], socket.socket
119 ] = get_port_socket,
120 **kwargs: Any,
121 ) -> None:
122 self._loop = loop
123 self.runner: Optional[BaseRunner] = None
124 self._root: Optional[URL] = None
125 self.host = host
126 self.port = port
127 self._closed = False
128 self.scheme = scheme
129 self.skip_url_asserts = skip_url_asserts
130 self.socket_factory = socket_factory
131
132 async def start_server(
133 self, loop: Optional[asyncio.AbstractEventLoop] = None, **kwargs: Any
134 ) -> None:
135 if self.runner:
136 return
137 self._loop = loop
138 self._ssl = kwargs.pop("ssl", None)
139 self.runner = await self._make_runner(handler_cancellation=True, **kwargs)
140 await self.runner.setup()
141 if not self.port:
142 self.port = 0
143 absolute_host = self.host
144 try:
145 version = ipaddress.ip_address(self.host).version
146 except ValueError:
147 version = 4
148 if version == 6:
149 absolute_host = f"[{self.host}]"
150 family = socket.AF_INET6 if version == 6 else socket.AF_INET
151 _sock = self.socket_factory(self.host, self.port, family)
152 self.host, self.port = _sock.getsockname()[:2]
153 site = SockSite(self.runner, sock=_sock, ssl_context=self._ssl)
154 await site.start()
155 server = site._server
156 assert server is not None
157 sockets = server.sockets # type: ignore[attr-defined]
158 assert sockets is not None
159 self.port = sockets[0].getsockname()[1]
160 if not self.scheme:
161 self.scheme = "https" if self._ssl else "http"
162 self._root = URL(f"{self.scheme}://{absolute_host}:{self.port}")
163
164 @abstractmethod # pragma: no cover
165 async def _make_runner(self, **kwargs: Any) -> BaseRunner:
166 pass
167
168 def make_url(self, path: StrOrURL) -> URL:
169 assert self._root is not None
170 url = URL(path)
171 if not self.skip_url_asserts:
172 assert not url.absolute
173 return self._root.join(url)
174 else:
175 return URL(str(self._root) + str(path))
176
177 @property
178 def started(self) -> bool:
179 return self.runner is not None
180
181 @property
182 def closed(self) -> bool:
183 return self._closed
184
185 @property
186 def handler(self) -> Server:
187 # for backward compatibility
188 # web.Server instance
189 runner = self.runner
190 assert runner is not None
191 assert runner.server is not None
192 return runner.server
193
194 async def close(self) -> None:
195 """Close all fixtures created by the test client.
196
197 After that point, the TestClient is no longer usable.
198
199 This is an idempotent function: running close multiple times
200 will not have any additional effects.
201
202 close is also run when the object is garbage collected, and on
203 exit when used as a context manager.
204
205 """
206 if self.started and not self.closed:
207 assert self.runner is not None
208 await self.runner.cleanup()
209 self._root = None
210 self.port = None
211 self._closed = True
212
213 def __enter__(self) -> None:
214 raise TypeError("Use async with instead")
215
216 def __exit__(
217 self,
218 exc_type: Optional[Type[BaseException]],
219 exc_value: Optional[BaseException],
220 traceback: Optional[TracebackType],
221 ) -> None:
222 # __exit__ should exist in pair with __enter__ but never executed
223 pass # pragma: no cover
224
225 async def __aenter__(self) -> "BaseTestServer":
226 await self.start_server(loop=self._loop)
227 return self
228
229 async def __aexit__(
230 self,
231 exc_type: Optional[Type[BaseException]],
232 exc_value: Optional[BaseException],
233 traceback: Optional[TracebackType],
234 ) -> None:
235 await self.close()
236
237
238class TestServer(BaseTestServer):
239 def __init__(
240 self,
241 app: Application,
242 *,
243 scheme: str = "",
244 host: str = "127.0.0.1",
245 port: Optional[int] = None,
246 **kwargs: Any,
247 ):
248 self.app = app
249 super().__init__(scheme=scheme, host=host, port=port, **kwargs)
250
251 async def _make_runner(self, **kwargs: Any) -> BaseRunner:
252 return AppRunner(self.app, **kwargs)
253
254
255class RawTestServer(BaseTestServer):
256 def __init__(
257 self,
258 handler: _RequestHandler,
259 *,
260 scheme: str = "",
261 host: str = "127.0.0.1",
262 port: Optional[int] = None,
263 **kwargs: Any,
264 ) -> None:
265 self._handler = handler
266 super().__init__(scheme=scheme, host=host, port=port, **kwargs)
267
268 async def _make_runner(self, debug: bool = True, **kwargs: Any) -> ServerRunner:
269 srv = Server(self._handler, loop=self._loop, debug=debug, **kwargs)
270 return ServerRunner(srv, debug=debug, **kwargs)
271
272
273class TestClient(Generic[_Request, _ApplicationNone]):
274 """
275 A test client implementation.
276
277 To write functional tests for aiohttp based servers.
278
279 """
280
281 __test__ = False
282
283 @overload
284 def __init__(
285 self: "TestClient[Request, Application]",
286 server: TestServer,
287 *,
288 cookie_jar: Optional[AbstractCookieJar] = None,
289 **kwargs: Any,
290 ) -> None: ...
291 @overload
292 def __init__(
293 self: "TestClient[_Request, None]",
294 server: BaseTestServer,
295 *,
296 cookie_jar: Optional[AbstractCookieJar] = None,
297 **kwargs: Any,
298 ) -> None: ...
299 def __init__(
300 self,
301 server: BaseTestServer,
302 *,
303 cookie_jar: Optional[AbstractCookieJar] = None,
304 loop: Optional[asyncio.AbstractEventLoop] = None,
305 **kwargs: Any,
306 ) -> None:
307 if not isinstance(server, BaseTestServer):
308 raise TypeError(
309 "server must be TestServer instance, found type: %r" % type(server)
310 )
311 self._server = server
312 self._loop = loop
313 if cookie_jar is None:
314 cookie_jar = aiohttp.CookieJar(unsafe=True, loop=loop)
315 self._session = ClientSession(loop=loop, cookie_jar=cookie_jar, **kwargs)
316 self._session._retry_connection = False
317 self._closed = False
318 self._responses: List[ClientResponse] = []
319 self._websockets: List[ClientWebSocketResponse] = []
320
321 async def start_server(self) -> None:
322 await self._server.start_server(loop=self._loop)
323
324 @property
325 def host(self) -> str:
326 return self._server.host
327
328 @property
329 def port(self) -> Optional[int]:
330 return self._server.port
331
332 @property
333 def server(self) -> BaseTestServer:
334 return self._server
335
336 @property
337 def app(self) -> _ApplicationNone:
338 return getattr(self._server, "app", None) # type: ignore[return-value]
339
340 @property
341 def session(self) -> ClientSession:
342 """An internal aiohttp.ClientSession.
343
344 Unlike the methods on the TestClient, client session requests
345 do not automatically include the host in the url queried, and
346 will require an absolute path to the resource.
347
348 """
349 return self._session
350
351 def make_url(self, path: StrOrURL) -> URL:
352 return self._server.make_url(path)
353
354 async def _request(
355 self, method: str, path: StrOrURL, **kwargs: Any
356 ) -> ClientResponse:
357 resp = await self._session.request(method, self.make_url(path), **kwargs)
358 # save it to close later
359 self._responses.append(resp)
360 return resp
361
362 if sys.version_info >= (3, 11) and TYPE_CHECKING:
363
364 def request(
365 self, method: str, path: StrOrURL, **kwargs: Unpack[_RequestOptions]
366 ) -> _RequestContextManager: ...
367
368 def get(
369 self,
370 path: StrOrURL,
371 **kwargs: Unpack[_RequestOptions],
372 ) -> _RequestContextManager: ...
373
374 def options(
375 self,
376 path: StrOrURL,
377 **kwargs: Unpack[_RequestOptions],
378 ) -> _RequestContextManager: ...
379
380 def head(
381 self,
382 path: StrOrURL,
383 **kwargs: Unpack[_RequestOptions],
384 ) -> _RequestContextManager: ...
385
386 def post(
387 self,
388 path: StrOrURL,
389 **kwargs: Unpack[_RequestOptions],
390 ) -> _RequestContextManager: ...
391
392 def put(
393 self,
394 path: StrOrURL,
395 **kwargs: Unpack[_RequestOptions],
396 ) -> _RequestContextManager: ...
397
398 def patch(
399 self,
400 path: StrOrURL,
401 **kwargs: Unpack[_RequestOptions],
402 ) -> _RequestContextManager: ...
403
404 def delete(
405 self,
406 path: StrOrURL,
407 **kwargs: Unpack[_RequestOptions],
408 ) -> _RequestContextManager: ...
409
410 else:
411
412 def request(
413 self, method: str, path: StrOrURL, **kwargs: Any
414 ) -> _RequestContextManager:
415 """Routes a request to tested http server.
416
417 The interface is identical to aiohttp.ClientSession.request,
418 except the loop kwarg is overridden by the instance used by the
419 test server.
420
421 """
422 return _RequestContextManager(self._request(method, path, **kwargs))
423
424 def get(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
425 """Perform an HTTP GET request."""
426 return _RequestContextManager(self._request(hdrs.METH_GET, path, **kwargs))
427
428 def post(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
429 """Perform an HTTP POST request."""
430 return _RequestContextManager(self._request(hdrs.METH_POST, path, **kwargs))
431
432 def options(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
433 """Perform an HTTP OPTIONS request."""
434 return _RequestContextManager(
435 self._request(hdrs.METH_OPTIONS, path, **kwargs)
436 )
437
438 def head(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
439 """Perform an HTTP HEAD request."""
440 return _RequestContextManager(self._request(hdrs.METH_HEAD, path, **kwargs))
441
442 def put(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
443 """Perform an HTTP PUT request."""
444 return _RequestContextManager(self._request(hdrs.METH_PUT, path, **kwargs))
445
446 def patch(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
447 """Perform an HTTP PATCH request."""
448 return _RequestContextManager(
449 self._request(hdrs.METH_PATCH, path, **kwargs)
450 )
451
452 def delete(self, path: StrOrURL, **kwargs: Any) -> _RequestContextManager:
453 """Perform an HTTP PATCH request."""
454 return _RequestContextManager(
455 self._request(hdrs.METH_DELETE, path, **kwargs)
456 )
457
458 def ws_connect(self, path: StrOrURL, **kwargs: Any) -> _WSRequestContextManager:
459 """Initiate websocket connection.
460
461 The api corresponds to aiohttp.ClientSession.ws_connect.
462
463 """
464 return _WSRequestContextManager(self._ws_connect(path, **kwargs))
465
466 async def _ws_connect(
467 self, path: StrOrURL, **kwargs: Any
468 ) -> ClientWebSocketResponse:
469 ws = await self._session.ws_connect(self.make_url(path), **kwargs)
470 self._websockets.append(ws)
471 return ws
472
473 async def close(self) -> None:
474 """Close all fixtures created by the test client.
475
476 After that point, the TestClient is no longer usable.
477
478 This is an idempotent function: running close multiple times
479 will not have any additional effects.
480
481 close is also run on exit when used as a(n) (asynchronous)
482 context manager.
483
484 """
485 if not self._closed:
486 for resp in self._responses:
487 resp.close()
488 for ws in self._websockets:
489 await ws.close()
490 await self._session.close()
491 await self._server.close()
492 self._closed = True
493
494 def __enter__(self) -> None:
495 raise TypeError("Use async with instead")
496
497 def __exit__(
498 self,
499 exc_type: Optional[Type[BaseException]],
500 exc: Optional[BaseException],
501 tb: Optional[TracebackType],
502 ) -> None:
503 # __exit__ should exist in pair with __enter__ but never executed
504 pass # pragma: no cover
505
506 async def __aenter__(self) -> Self:
507 await self.start_server()
508 return self
509
510 async def __aexit__(
511 self,
512 exc_type: Optional[Type[BaseException]],
513 exc: Optional[BaseException],
514 tb: Optional[TracebackType],
515 ) -> None:
516 await self.close()
517
518
519class AioHTTPTestCase(IsolatedAsyncioTestCase):
520 """A base class to allow for unittest web applications using aiohttp.
521
522 Provides the following:
523
524 * self.client (aiohttp.test_utils.TestClient): an aiohttp test client.
525 * self.loop (asyncio.BaseEventLoop): the event loop in which the
526 application and server are running.
527 * self.app (aiohttp.web.Application): the application returned by
528 self.get_application()
529
530 Note that the TestClient's methods are asynchronous: you have to
531 execute function on the test client using asynchronous methods.
532 """
533
534 async def get_application(self) -> Application:
535 """Get application.
536
537 This method should be overridden
538 to return the aiohttp.web.Application
539 object to test.
540 """
541 return self.get_app()
542
543 def get_app(self) -> Application:
544 """Obsolete method used to constructing web application.
545
546 Use .get_application() coroutine instead.
547 """
548 raise RuntimeError("Did you forget to define get_application()?")
549
550 async def asyncSetUp(self) -> None:
551 self.loop = asyncio.get_running_loop()
552 return await self.setUpAsync()
553
554 async def setUpAsync(self) -> None:
555 self.app = await self.get_application()
556 self.server = await self.get_server(self.app)
557 self.client = await self.get_client(self.server)
558
559 await self.client.start_server()
560
561 async def asyncTearDown(self) -> None:
562 return await self.tearDownAsync()
563
564 async def tearDownAsync(self) -> None:
565 await self.client.close()
566
567 async def get_server(self, app: Application) -> TestServer:
568 """Return a TestServer instance."""
569 return TestServer(app, loop=self.loop)
570
571 async def get_client(self, server: TestServer) -> TestClient[Request, Application]:
572 """Return a TestClient instance."""
573 return TestClient(server, loop=self.loop)
574
575
576def unittest_run_loop(func: Any, *args: Any, **kwargs: Any) -> Any:
577 """
578 A decorator dedicated to use with asynchronous AioHTTPTestCase test methods.
579
580 In 3.8+, this does nothing.
581 """
582 warnings.warn(
583 "Decorator `@unittest_run_loop` is no longer needed in aiohttp 3.8+",
584 DeprecationWarning,
585 stacklevel=2,
586 )
587 return func
588
589
590_LOOP_FACTORY = Callable[[], asyncio.AbstractEventLoop]
591
592
593@contextlib.contextmanager
594def loop_context(
595 loop_factory: _LOOP_FACTORY = asyncio.new_event_loop, fast: bool = False
596) -> Iterator[asyncio.AbstractEventLoop]:
597 """A contextmanager that creates an event_loop, for test purposes.
598
599 Handles the creation and cleanup of a test loop.
600 """
601 loop = setup_test_loop(loop_factory)
602 yield loop
603 teardown_test_loop(loop, fast=fast)
604
605
606def setup_test_loop(
607 loop_factory: _LOOP_FACTORY = asyncio.new_event_loop,
608) -> asyncio.AbstractEventLoop:
609 """Create and return an asyncio.BaseEventLoop instance.
610
611 The caller should also call teardown_test_loop,
612 once they are done with the loop.
613 """
614 loop = loop_factory()
615 asyncio.set_event_loop(loop)
616 return loop
617
618
619def teardown_test_loop(loop: asyncio.AbstractEventLoop, fast: bool = False) -> None:
620 """Teardown and cleanup an event_loop created by setup_test_loop."""
621 closed = loop.is_closed()
622 if not closed:
623 loop.call_soon(loop.stop)
624 loop.run_forever()
625 loop.close()
626
627 if not fast:
628 gc.collect()
629
630 asyncio.set_event_loop(None)
631
632
633def _create_app_mock() -> mock.MagicMock:
634 def get_dict(app: Any, key: str) -> Any:
635 return app.__app_dict[key]
636
637 def set_dict(app: Any, key: str, value: Any) -> None:
638 app.__app_dict[key] = value
639
640 app = mock.MagicMock(spec=Application)
641 app.__app_dict = {}
642 app.__getitem__ = get_dict
643 app.__setitem__ = set_dict
644
645 app._debug = False
646 app.on_response_prepare = Signal(app)
647 app.on_response_prepare.freeze()
648 return app
649
650
651def _create_transport(sslcontext: Optional[SSLContext] = None) -> mock.Mock:
652 transport = mock.Mock()
653
654 def get_extra_info(key: str) -> Optional[SSLContext]:
655 if key == "sslcontext":
656 return sslcontext
657 else:
658 return None
659
660 transport.get_extra_info.side_effect = get_extra_info
661 return transport
662
663
664def make_mocked_request(
665 method: str,
666 path: str,
667 headers: Any = None,
668 *,
669 match_info: Any = sentinel,
670 version: HttpVersion = HttpVersion(1, 1),
671 closing: bool = False,
672 app: Any = None,
673 writer: Any = sentinel,
674 protocol: Any = sentinel,
675 transport: Any = sentinel,
676 payload: StreamReader = EMPTY_PAYLOAD,
677 sslcontext: Optional[SSLContext] = None,
678 client_max_size: int = 1024**2,
679 loop: Any = ...,
680) -> Request:
681 """Creates mocked web.Request testing purposes.
682
683 Useful in unit tests, when spinning full web server is overkill or
684 specific conditions and errors are hard to trigger.
685 """
686 task = mock.Mock()
687 if loop is ...:
688 # no loop passed, try to get the current one if
689 # its is running as we need a real loop to create
690 # executor jobs to be able to do testing
691 # with a real executor
692 try:
693 loop = asyncio.get_running_loop()
694 except RuntimeError:
695 loop = mock.Mock()
696 loop.create_future.return_value = ()
697
698 if version < HttpVersion(1, 1):
699 closing = True
700
701 if headers:
702 headers = CIMultiDictProxy(CIMultiDict(headers))
703 raw_hdrs = tuple(
704 (k.encode("utf-8"), v.encode("utf-8")) for k, v in headers.items()
705 )
706 else:
707 headers = CIMultiDictProxy(CIMultiDict())
708 raw_hdrs = ()
709
710 chunked = "chunked" in headers.get(hdrs.TRANSFER_ENCODING, "").lower()
711
712 message = RawRequestMessage(
713 method,
714 path,
715 version,
716 headers,
717 raw_hdrs,
718 closing,
719 None,
720 False,
721 chunked,
722 URL(path),
723 )
724 if app is None:
725 app = _create_app_mock()
726
727 if transport is sentinel:
728 transport = _create_transport(sslcontext)
729
730 if protocol is sentinel:
731 protocol = mock.Mock()
732 protocol.max_field_size = 8190
733 protocol.max_line_length = 8190
734 protocol.max_headers = 128
735 protocol.transport = transport
736 type(protocol).peername = mock.PropertyMock(
737 return_value=transport.get_extra_info("peername")
738 )
739 type(protocol).ssl_context = mock.PropertyMock(return_value=sslcontext)
740
741 if writer is sentinel:
742 writer = mock.Mock()
743 writer.write_headers = make_mocked_coro(None)
744 writer.write = make_mocked_coro(None)
745 writer.write_eof = make_mocked_coro(None)
746 writer.drain = make_mocked_coro(None)
747 writer.transport = transport
748
749 protocol.transport = transport
750 protocol.writer = writer
751
752 req = Request(
753 message, payload, protocol, writer, task, loop, client_max_size=client_max_size
754 )
755
756 match_info = UrlMappingMatchInfo(
757 {} if match_info is sentinel else match_info, mock.Mock()
758 )
759 match_info.add_app(app)
760 req._match_info = match_info
761
762 return req
763
764
765def make_mocked_coro(
766 return_value: Any = sentinel, raise_exception: Any = sentinel
767) -> Any:
768 """Creates a coroutine mock."""
769
770 async def mock_coro(*args: Any, **kwargs: Any) -> Any:
771 if raise_exception is not sentinel:
772 raise raise_exception
773 if not inspect.isawaitable(return_value):
774 return return_value
775 await return_value
776
777 return mock.Mock(wraps=mock_coro)
778 