Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
web_runner.py400 linesDownload Raw Back to aiohttp
1import asyncio
2import signal
3import socket
4import warnings
5from abc import ABC, abstractmethod
6from typing import TYPE_CHECKING, Any, List, Optional, Set
7
8from yarl import URL
9
10from .typedefs import PathLike
11from .web_app import Application
12from .web_server import Server
13
14if TYPE_CHECKING:
15    from ssl import SSLContext
16else:
17    try:
18        from ssl import SSLContext
19    except ImportError:  # pragma: no cover
20        SSLContext = object  # type: ignore[misc,assignment]
21
22__all__ = (
23    "BaseSite",
24    "TCPSite",
25    "UnixSite",
26    "NamedPipeSite",
27    "SockSite",
28    "BaseRunner",
29    "AppRunner",
30    "ServerRunner",
31    "GracefulExit",
32)
33
34
35class GracefulExit(SystemExit):
36    code = 1
37
38
39def _raise_graceful_exit() -> None:
40    raise GracefulExit()
41
42
43class BaseSite(ABC):
44    __slots__ = ("_runner", "_ssl_context", "_backlog", "_server")
45
46    def __init__(
47        self,
48        runner: "BaseRunner",
49        *,
50        shutdown_timeout: float = 60.0,
51        ssl_context: Optional[SSLContext] = None,
52        backlog: int = 128,
53    ) -> None:
54        if runner.server is None:
55            raise RuntimeError("Call runner.setup() before making a site")
56        if shutdown_timeout != 60.0:
57            msg = "shutdown_timeout should be set on BaseRunner"
58            warnings.warn(msg, DeprecationWarning, stacklevel=2)
59            runner._shutdown_timeout = shutdown_timeout
60        self._runner = runner
61        self._ssl_context = ssl_context
62        self._backlog = backlog
63        self._server: Optional[asyncio.AbstractServer] = None
64
65    @property
66    @abstractmethod
67    def name(self) -> str:
68        pass  # pragma: no cover
69
70    @abstractmethod
71    async def start(self) -> None:
72        self._runner._reg_site(self)
73
74    async def stop(self) -> None:
75        self._runner._check_site(self)
76        if self._server is not None:  # Maybe not started yet
77            self._server.close()
78
79        self._runner._unreg_site(self)
80
81
82class TCPSite(BaseSite):
83    __slots__ = ("_host", "_port", "_reuse_address", "_reuse_port")
84
85    def __init__(
86        self,
87        runner: "BaseRunner",
88        host: Optional[str] = None,
89        port: Optional[int] = None,
90        *,
91        shutdown_timeout: float = 60.0,
92        ssl_context: Optional[SSLContext] = None,
93        backlog: int = 128,
94        reuse_address: Optional[bool] = None,
95        reuse_port: Optional[bool] = None,
96    ) -> None:
97        super().__init__(
98            runner,
99            shutdown_timeout=shutdown_timeout,
100            ssl_context=ssl_context,
101            backlog=backlog,
102        )
103        self._host = host
104        if port is None:
105            port = 8443 if self._ssl_context else 8080
106        self._port = port
107        self._reuse_address = reuse_address
108        self._reuse_port = reuse_port
109
110    @property
111    def name(self) -> str:
112        scheme = "https" if self._ssl_context else "http"
113        host = "0.0.0.0" if not self._host else self._host
114        return str(URL.build(scheme=scheme, host=host, port=self._port))
115
116    async def start(self) -> None:
117        await super().start()
118        loop = asyncio.get_event_loop()
119        server = self._runner.server
120        assert server is not None
121        self._server = await loop.create_server(
122            server,
123            self._host,
124            self._port,
125            ssl=self._ssl_context,
126            backlog=self._backlog,
127            reuse_address=self._reuse_address,
128            reuse_port=self._reuse_port,
129        )
130
131
132class UnixSite(BaseSite):
133    __slots__ = ("_path",)
134
135    def __init__(
136        self,
137        runner: "BaseRunner",
138        path: PathLike,
139        *,
140        shutdown_timeout: float = 60.0,
141        ssl_context: Optional[SSLContext] = None,
142        backlog: int = 128,
143    ) -> None:
144        super().__init__(
145            runner,
146            shutdown_timeout=shutdown_timeout,
147            ssl_context=ssl_context,
148            backlog=backlog,
149        )
150        self._path = path
151
152    @property
153    def name(self) -> str:
154        scheme = "https" if self._ssl_context else "http"
155        return f"{scheme}://unix:{self._path}:"
156
157    async def start(self) -> None:
158        await super().start()
159        loop = asyncio.get_event_loop()
160        server = self._runner.server
161        assert server is not None
162        self._server = await loop.create_unix_server(
163            server,
164            self._path,
165            ssl=self._ssl_context,
166            backlog=self._backlog,
167        )
168
169
170class NamedPipeSite(BaseSite):
171    __slots__ = ("_path",)
172
173    def __init__(
174        self, runner: "BaseRunner", path: str, *, shutdown_timeout: float = 60.0
175    ) -> None:
176        loop = asyncio.get_event_loop()
177        if not isinstance(
178            loop, asyncio.ProactorEventLoop  # type: ignore[attr-defined]
179        ):
180            raise RuntimeError(
181                "Named Pipes only available in proactor loop under windows"
182            )
183        super().__init__(runner, shutdown_timeout=shutdown_timeout)
184        self._path = path
185
186    @property
187    def name(self) -> str:
188        return self._path
189
190    async def start(self) -> None:
191        await super().start()
192        loop = asyncio.get_event_loop()
193        server = self._runner.server
194        assert server is not None
195        _server = await loop.start_serving_pipe(  # type: ignore[attr-defined]
196            server, self._path
197        )
198        self._server = _server[0]
199
200
201class SockSite(BaseSite):
202    __slots__ = ("_sock", "_name")
203
204    def __init__(
205        self,
206        runner: "BaseRunner",
207        sock: socket.socket,
208        *,
209        shutdown_timeout: float = 60.0,
210        ssl_context: Optional[SSLContext] = None,
211        backlog: int = 128,
212    ) -> None:
213        super().__init__(
214            runner,
215            shutdown_timeout=shutdown_timeout,
216            ssl_context=ssl_context,
217            backlog=backlog,
218        )
219        self._sock = sock
220        scheme = "https" if self._ssl_context else "http"
221        if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
222            name = f"{scheme}://unix:{sock.getsockname()}:"
223        else:
224            host, port = sock.getsockname()[:2]
225            name = str(URL.build(scheme=scheme, host=host, port=port))
226        self._name = name
227
228    @property
229    def name(self) -> str:
230        return self._name
231
232    async def start(self) -> None:
233        await super().start()
234        loop = asyncio.get_event_loop()
235        server = self._runner.server
236        assert server is not None
237        self._server = await loop.create_server(
238            server, sock=self._sock, ssl=self._ssl_context, backlog=self._backlog
239        )
240
241
242class BaseRunner(ABC):
243    __slots__ = ("_handle_signals", "_kwargs", "_server", "_sites", "_shutdown_timeout")
244
245    def __init__(
246        self,
247        *,
248        handle_signals: bool = False,
249        shutdown_timeout: float = 60.0,
250        **kwargs: Any,
251    ) -> None:
252        self._handle_signals = handle_signals
253        self._kwargs = kwargs
254        self._server: Optional[Server] = None
255        self._sites: List[BaseSite] = []
256        self._shutdown_timeout = shutdown_timeout
257
258    @property
259    def server(self) -> Optional[Server]:
260        return self._server
261
262    @property
263    def addresses(self) -> List[Any]:
264        ret: List[Any] = []
265        for site in self._sites:
266            server = site._server
267            if server is not None:
268                sockets = server.sockets  # type: ignore[attr-defined]
269                if sockets is not None:
270                    for sock in sockets:
271                        ret.append(sock.getsockname())
272        return ret
273
274    @property
275    def sites(self) -> Set[BaseSite]:
276        return set(self._sites)
277
278    async def setup(self) -> None:
279        loop = asyncio.get_event_loop()
280
281        if self._handle_signals:
282            try:
283                loop.add_signal_handler(signal.SIGINT, _raise_graceful_exit)
284                loop.add_signal_handler(signal.SIGTERM, _raise_graceful_exit)
285            except NotImplementedError:  # pragma: no cover
286                # add_signal_handler is not implemented on Windows
287                pass
288
289        self._server = await self._make_server()
290
291    @abstractmethod
292    async def shutdown(self) -> None:
293        """Call any shutdown hooks to help server close gracefully."""
294
295    async def cleanup(self) -> None:
296        # The loop over sites is intentional, an exception on gather()
297        # leaves self._sites in unpredictable state.
298        # The loop guaranties that a site is either deleted on success or
299        # still present on failure
300        for site in list(self._sites):
301            await site.stop()
302
303        if self._server:  # If setup succeeded
304            # Yield to event loop to ensure incoming requests prior to stopping the sites
305            # have all started to be handled before we proceed to close idle connections.
306            await asyncio.sleep(0)
307            self._server.pre_shutdown()
308            await self.shutdown()
309            await self._server.shutdown(self._shutdown_timeout)
310        await self._cleanup_server()
311
312        self._server = None
313        if self._handle_signals:
314            loop = asyncio.get_running_loop()
315            try:
316                loop.remove_signal_handler(signal.SIGINT)
317                loop.remove_signal_handler(signal.SIGTERM)
318            except NotImplementedError:  # pragma: no cover
319                # remove_signal_handler is not implemented on Windows
320                pass
321
322    @abstractmethod
323    async def _make_server(self) -> Server:
324        pass  # pragma: no cover
325
326    @abstractmethod
327    async def _cleanup_server(self) -> None:
328        pass  # pragma: no cover
329
330    def _reg_site(self, site: BaseSite) -> None:
331        if site in self._sites:
332            raise RuntimeError(f"Site {site} is already registered in runner {self}")
333        self._sites.append(site)
334
335    def _check_site(self, site: BaseSite) -> None:
336        if site not in self._sites:
337            raise RuntimeError(f"Site {site} is not registered in runner {self}")
338
339    def _unreg_site(self, site: BaseSite) -> None:
340        if site not in self._sites:
341            raise RuntimeError(f"Site {site} is not registered in runner {self}")
342        self._sites.remove(site)
343
344
345class ServerRunner(BaseRunner):
346    """Low-level web server runner"""
347
348    __slots__ = ("_web_server",)
349
350    def __init__(
351        self, web_server: Server, *, handle_signals: bool = False, **kwargs: Any
352    ) -> None:
353        super().__init__(handle_signals=handle_signals, **kwargs)
354        self._web_server = web_server
355
356    async def shutdown(self) -> None:
357        pass
358
359    async def _make_server(self) -> Server:
360        return self._web_server
361
362    async def _cleanup_server(self) -> None:
363        pass
364
365
366class AppRunner(BaseRunner):
367    """Web Application runner"""
368
369    __slots__ = ("_app",)
370
371    def __init__(
372        self, app: Application, *, handle_signals: bool = False, **kwargs: Any
373    ) -> None:
374        super().__init__(handle_signals=handle_signals, **kwargs)
375        if not isinstance(app, Application):
376            raise TypeError(
377                "The first argument should be web.Application "
378                "instance, got {!r}".format(app)
379            )
380        self._app = app
381
382    @property
383    def app(self) -> Application:
384        return self._app
385
386    async def shutdown(self) -> None:
387        await self._app.shutdown()
388
389    async def _make_server(self) -> Server:
390        loop = asyncio.get_event_loop()
391        self._app._set_loop(loop)
392        self._app.on_startup.freeze()
393        await self._app.startup()
394        self._app.freeze()
395
396        return self._app._make_handler(loop=loop, **self._kwargs)
397
398    async def _cleanup_server(self) -> None:
399        await self._app.cleanup()
400 
codekingpro/portable-devtools · Team Ai