Coverage for /pythoncovmergedfiles/medio/medio/src/aiohttp/aiohttp/web_runner.py: 39%
231 statements
« prev ^ index » next coverage.py v7.3.1, created at 2023-09-27 06:09 +0000
« prev ^ index » next coverage.py v7.3.1, created at 2023-09-27 06:09 +0000
1import asyncio
2import signal
3import socket
4from abc import ABC, abstractmethod
5from contextlib import suppress
6from typing import Any, List, Optional, Set, Type
8from yarl import URL
10from .abc import AbstractAccessLogger, AbstractStreamWriter
11from .http_parser import RawRequestMessage
12from .streams import StreamReader
13from .typedefs import PathLike
14from .web_app import Application
15from .web_log import AccessLogger
16from .web_protocol import RequestHandler
17from .web_request import Request
18from .web_server import Server
20try:
21 from ssl import SSLContext
22except ImportError:
23 SSLContext = object # type: ignore[misc,assignment]
26__all__ = (
27 "BaseSite",
28 "TCPSite",
29 "UnixSite",
30 "NamedPipeSite",
31 "SockSite",
32 "BaseRunner",
33 "AppRunner",
34 "ServerRunner",
35 "GracefulExit",
36)
39class GracefulExit(SystemExit):
40 code = 1
43def _raise_graceful_exit() -> None:
44 raise GracefulExit()
47class BaseSite(ABC):
48 __slots__ = ("_runner", "_shutdown_timeout", "_ssl_context", "_backlog", "_server")
50 def __init__(
51 self,
52 runner: "BaseRunner",
53 *,
54 shutdown_timeout: float = 60.0,
55 ssl_context: Optional[SSLContext] = None,
56 backlog: int = 128,
57 ) -> None:
58 if runner.server is None:
59 raise RuntimeError("Call runner.setup() before making a site")
60 self._runner = runner
61 self._shutdown_timeout = shutdown_timeout
62 self._ssl_context = ssl_context
63 self._backlog = backlog
64 self._server: Optional[asyncio.AbstractServer] = None
66 @property
67 @abstractmethod
68 def name(self) -> str:
69 pass # pragma: no cover
71 @abstractmethod
72 async def start(self) -> None:
73 self._runner._reg_site(self)
75 async def stop(self) -> None:
76 self._runner._check_site(self)
77 if self._server is None:
78 self._runner._unreg_site(self)
79 return # not started yet
80 self._server.close()
81 # named pipes do not have wait_closed property
82 if hasattr(self._server, "wait_closed"):
83 await self._server.wait_closed()
85 # Wait for pending tasks for a given time limit.
86 with suppress(asyncio.TimeoutError):
87 await asyncio.wait_for(
88 self._wait(asyncio.current_task()), timeout=self._shutdown_timeout
89 )
91 await self._runner.shutdown()
92 assert self._runner.server
93 await self._runner.server.shutdown(self._shutdown_timeout)
94 self._runner._unreg_site(self)
96 async def _wait(self, parent_task: Optional["asyncio.Task[object]"]) -> None:
97 exclude = self._runner.starting_tasks | {asyncio.current_task(), parent_task}
98 while tasks := asyncio.all_tasks() - exclude:
99 await asyncio.wait(tasks)
102class TCPSite(BaseSite):
103 __slots__ = ("_host", "_port", "_reuse_address", "_reuse_port")
105 def __init__(
106 self,
107 runner: "BaseRunner",
108 host: Optional[str] = None,
109 port: Optional[int] = None,
110 *,
111 shutdown_timeout: float = 60.0,
112 ssl_context: Optional[SSLContext] = None,
113 backlog: int = 128,
114 reuse_address: Optional[bool] = None,
115 reuse_port: Optional[bool] = None,
116 ) -> None:
117 super().__init__(
118 runner,
119 shutdown_timeout=shutdown_timeout,
120 ssl_context=ssl_context,
121 backlog=backlog,
122 )
123 self._host = host
124 if port is None:
125 port = 8443 if self._ssl_context else 8080
126 self._port = port
127 self._reuse_address = reuse_address
128 self._reuse_port = reuse_port
130 @property
131 def name(self) -> str:
132 scheme = "https" if self._ssl_context else "http"
133 host = "0.0.0.0" if self._host is None else self._host
134 return str(URL.build(scheme=scheme, host=host, port=self._port))
136 async def start(self) -> None:
137 await super().start()
138 loop = asyncio.get_event_loop()
139 server = self._runner.server
140 assert server is not None
141 self._server = await loop.create_server(
142 server,
143 self._host,
144 self._port,
145 ssl=self._ssl_context,
146 backlog=self._backlog,
147 reuse_address=self._reuse_address,
148 reuse_port=self._reuse_port,
149 )
152class UnixSite(BaseSite):
153 __slots__ = ("_path",)
155 def __init__(
156 self,
157 runner: "BaseRunner",
158 path: PathLike,
159 *,
160 shutdown_timeout: float = 60.0,
161 ssl_context: Optional[SSLContext] = None,
162 backlog: int = 128,
163 ) -> None:
164 super().__init__(
165 runner,
166 shutdown_timeout=shutdown_timeout,
167 ssl_context=ssl_context,
168 backlog=backlog,
169 )
170 self._path = path
172 @property
173 def name(self) -> str:
174 scheme = "https" if self._ssl_context else "http"
175 return f"{scheme}://unix:{self._path}:"
177 async def start(self) -> None:
178 await super().start()
179 loop = asyncio.get_event_loop()
180 server = self._runner.server
181 assert server is not None
182 self._server = await loop.create_unix_server(
183 server,
184 self._path,
185 ssl=self._ssl_context,
186 backlog=self._backlog,
187 )
190class NamedPipeSite(BaseSite):
191 __slots__ = ("_path",)
193 def __init__(
194 self, runner: "BaseRunner", path: str, *, shutdown_timeout: float = 60.0
195 ) -> None:
196 loop = asyncio.get_event_loop()
197 if not isinstance(
198 loop, asyncio.ProactorEventLoop # type: ignore[attr-defined]
199 ):
200 raise RuntimeError(
201 "Named Pipes only available in proactor" "loop under windows"
202 )
203 super().__init__(runner, shutdown_timeout=shutdown_timeout)
204 self._path = path
206 @property
207 def name(self) -> str:
208 return self._path
210 async def start(self) -> None:
211 await super().start()
212 loop = asyncio.get_event_loop()
213 server = self._runner.server
214 assert server is not None
215 _server = await loop.start_serving_pipe( # type: ignore[attr-defined]
216 server, self._path
217 )
218 self._server = _server[0]
221class SockSite(BaseSite):
222 __slots__ = ("_sock", "_name")
224 def __init__(
225 self,
226 runner: "BaseRunner",
227 sock: socket.socket,
228 *,
229 shutdown_timeout: float = 60.0,
230 ssl_context: Optional[SSLContext] = None,
231 backlog: int = 128,
232 ) -> None:
233 super().__init__(
234 runner,
235 shutdown_timeout=shutdown_timeout,
236 ssl_context=ssl_context,
237 backlog=backlog,
238 )
239 self._sock = sock
240 scheme = "https" if self._ssl_context else "http"
241 if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
242 name = f"{scheme}://unix:{sock.getsockname()}:"
243 else:
244 host, port = sock.getsockname()[:2]
245 name = str(URL.build(scheme=scheme, host=host, port=port))
246 self._name = name
248 @property
249 def name(self) -> str:
250 return self._name
252 async def start(self) -> None:
253 await super().start()
254 loop = asyncio.get_event_loop()
255 server = self._runner.server
256 assert server is not None
257 self._server = await loop.create_server(
258 server, sock=self._sock, ssl=self._ssl_context, backlog=self._backlog
259 )
262class BaseRunner(ABC):
263 __slots__ = ("starting_tasks", "_handle_signals", "_kwargs", "_server", "_sites")
265 def __init__(self, *, handle_signals: bool = False, **kwargs: Any) -> None:
266 self._handle_signals = handle_signals
267 self._kwargs = kwargs
268 self._server: Optional[Server] = None
269 self._sites: List[BaseSite] = []
271 @property
272 def server(self) -> Optional[Server]:
273 return self._server
275 @property
276 def addresses(self) -> List[Any]:
277 ret: List[Any] = []
278 for site in self._sites:
279 server = site._server
280 if server is not None:
281 sockets = server.sockets # type: ignore[attr-defined]
282 if sockets is not None:
283 for sock in sockets:
284 ret.append(sock.getsockname())
285 return ret
287 @property
288 def sites(self) -> Set[BaseSite]:
289 return set(self._sites)
291 async def setup(self) -> None:
292 loop = asyncio.get_event_loop()
294 if self._handle_signals:
295 try:
296 loop.add_signal_handler(signal.SIGINT, _raise_graceful_exit)
297 loop.add_signal_handler(signal.SIGTERM, _raise_graceful_exit)
298 except NotImplementedError: # pragma: no cover
299 # add_signal_handler is not implemented on Windows
300 pass
302 self._server = await self._make_server()
303 # On shutdown we want to avoid waiting on tasks which run forever.
304 # It's very likely that all tasks which run forever will have been created by
305 # the time we have completed the application startup (in self._make_server()),
306 # so we just record all running tasks here and exclude them later.
307 self.starting_tasks = asyncio.all_tasks()
309 @abstractmethod
310 async def shutdown(self) -> None:
311 pass # pragma: no cover
313 async def cleanup(self) -> None:
314 loop = asyncio.get_event_loop()
316 # The loop over sites is intentional, an exception on gather()
317 # leaves self._sites in unpredictable state.
318 # The loop guarantees that a site is either deleted on success or
319 # still present on failure
320 for site in list(self._sites):
321 await site.stop()
322 await self._cleanup_server()
323 self._server = None
324 if self._handle_signals:
325 try:
326 loop.remove_signal_handler(signal.SIGINT)
327 loop.remove_signal_handler(signal.SIGTERM)
328 except NotImplementedError: # pragma: no cover
329 # remove_signal_handler is not implemented on Windows
330 pass
332 @abstractmethod
333 async def _make_server(self) -> Server:
334 pass # pragma: no cover
336 @abstractmethod
337 async def _cleanup_server(self) -> None:
338 pass # pragma: no cover
340 def _reg_site(self, site: BaseSite) -> None:
341 if site in self._sites:
342 raise RuntimeError(f"Site {site} is already registered in runner {self}")
343 self._sites.append(site)
345 def _check_site(self, site: BaseSite) -> None:
346 if site not in self._sites:
347 raise RuntimeError(f"Site {site} is not registered in runner {self}")
349 def _unreg_site(self, site: BaseSite) -> None:
350 if site not in self._sites:
351 raise RuntimeError(f"Site {site} is not registered in runner {self}")
352 self._sites.remove(site)
355class ServerRunner(BaseRunner):
356 """Low-level web server runner"""
358 __slots__ = ("_web_server",)
360 def __init__(
361 self, web_server: Server, *, handle_signals: bool = False, **kwargs: Any
362 ) -> None:
363 super().__init__(handle_signals=handle_signals, **kwargs)
364 self._web_server = web_server
366 async def shutdown(self) -> None:
367 pass
369 async def _make_server(self) -> Server:
370 return self._web_server
372 async def _cleanup_server(self) -> None:
373 pass
376class AppRunner(BaseRunner):
377 """Web Application runner"""
379 __slots__ = ("_app",)
381 def __init__(
382 self,
383 app: Application,
384 *,
385 handle_signals: bool = False,
386 access_log_class: Type[AbstractAccessLogger] = AccessLogger,
387 **kwargs: Any,
388 ) -> None:
389 if not isinstance(app, Application):
390 raise TypeError(
391 "The first argument should be web.Application "
392 "instance, got {!r}".format(app)
393 )
394 kwargs["access_log_class"] = access_log_class
396 if app._handler_args:
397 for k, v in app._handler_args.items():
398 kwargs[k] = v
400 if not issubclass(kwargs["access_log_class"], AbstractAccessLogger):
401 raise TypeError(
402 "access_log_class must be subclass of "
403 "aiohttp.abc.AbstractAccessLogger, got {}".format(
404 kwargs["access_log_class"]
405 )
406 )
408 super().__init__(handle_signals=handle_signals, **kwargs)
409 self._app = app
411 @property
412 def app(self) -> Application:
413 return self._app
415 async def shutdown(self) -> None:
416 await self._app.shutdown()
418 async def _make_server(self) -> Server:
419 self._app.on_startup.freeze()
420 await self._app.startup()
421 self._app.freeze()
423 return Server(
424 self._app._handle, # type: ignore[arg-type]
425 request_factory=self._make_request,
426 **self._kwargs,
427 )
429 def _make_request(
430 self,
431 message: RawRequestMessage,
432 payload: StreamReader,
433 protocol: RequestHandler,
434 writer: AbstractStreamWriter,
435 task: "asyncio.Task[None]",
436 _cls: Type[Request] = Request,
437 ) -> Request:
438 loop = asyncio.get_running_loop()
439 return _cls(
440 message,
441 payload,
442 protocol,
443 writer,
444 task,
445 loop,
446 client_max_size=self.app._client_max_size,
447 )
449 async def _cleanup_server(self) -> None:
450 await self._app.cleanup()