Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
test_utils.py778 linesDownload Raw Back to aiohttp
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 
codekingpro/portable-devtools · Team Ai