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

1import asyncio 

2import signal 

3import socket 

4from abc import ABC, abstractmethod 

5from contextlib import suppress 

6from typing import Any, List, Optional, Set, Type 

7 

8from yarl import URL 

9 

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 

19 

20try: 

21 from ssl import SSLContext 

22except ImportError: 

23 SSLContext = object # type: ignore[misc,assignment] 

24 

25 

26__all__ = ( 

27 "BaseSite", 

28 "TCPSite", 

29 "UnixSite", 

30 "NamedPipeSite", 

31 "SockSite", 

32 "BaseRunner", 

33 "AppRunner", 

34 "ServerRunner", 

35 "GracefulExit", 

36) 

37 

38 

39class GracefulExit(SystemExit): 

40 code = 1 

41 

42 

43def _raise_graceful_exit() -> None: 

44 raise GracefulExit() 

45 

46 

47class BaseSite(ABC): 

48 __slots__ = ("_runner", "_shutdown_timeout", "_ssl_context", "_backlog", "_server") 

49 

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 

65 

66 @property 

67 @abstractmethod 

68 def name(self) -> str: 

69 pass # pragma: no cover 

70 

71 @abstractmethod 

72 async def start(self) -> None: 

73 self._runner._reg_site(self) 

74 

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() 

84 

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 ) 

90 

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) 

95 

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) 

100 

101 

102class TCPSite(BaseSite): 

103 __slots__ = ("_host", "_port", "_reuse_address", "_reuse_port") 

104 

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 

129 

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)) 

135 

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 ) 

150 

151 

152class UnixSite(BaseSite): 

153 __slots__ = ("_path",) 

154 

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 

171 

172 @property 

173 def name(self) -> str: 

174 scheme = "https" if self._ssl_context else "http" 

175 return f"{scheme}://unix:{self._path}:" 

176 

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 ) 

188 

189 

190class NamedPipeSite(BaseSite): 

191 __slots__ = ("_path",) 

192 

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 

205 

206 @property 

207 def name(self) -> str: 

208 return self._path 

209 

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] 

219 

220 

221class SockSite(BaseSite): 

222 __slots__ = ("_sock", "_name") 

223 

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 

247 

248 @property 

249 def name(self) -> str: 

250 return self._name 

251 

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 ) 

260 

261 

262class BaseRunner(ABC): 

263 __slots__ = ("starting_tasks", "_handle_signals", "_kwargs", "_server", "_sites") 

264 

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] = [] 

270 

271 @property 

272 def server(self) -> Optional[Server]: 

273 return self._server 

274 

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 

286 

287 @property 

288 def sites(self) -> Set[BaseSite]: 

289 return set(self._sites) 

290 

291 async def setup(self) -> None: 

292 loop = asyncio.get_event_loop() 

293 

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 

301 

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() 

308 

309 @abstractmethod 

310 async def shutdown(self) -> None: 

311 pass # pragma: no cover 

312 

313 async def cleanup(self) -> None: 

314 loop = asyncio.get_event_loop() 

315 

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 

331 

332 @abstractmethod 

333 async def _make_server(self) -> Server: 

334 pass # pragma: no cover 

335 

336 @abstractmethod 

337 async def _cleanup_server(self) -> None: 

338 pass # pragma: no cover 

339 

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) 

344 

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}") 

348 

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) 

353 

354 

355class ServerRunner(BaseRunner): 

356 """Low-level web server runner""" 

357 

358 __slots__ = ("_web_server",) 

359 

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 

365 

366 async def shutdown(self) -> None: 

367 pass 

368 

369 async def _make_server(self) -> Server: 

370 return self._web_server 

371 

372 async def _cleanup_server(self) -> None: 

373 pass 

374 

375 

376class AppRunner(BaseRunner): 

377 """Web Application runner""" 

378 

379 __slots__ = ("_app",) 

380 

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 

395 

396 if app._handler_args: 

397 for k, v in app._handler_args.items(): 

398 kwargs[k] = v 

399 

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 ) 

407 

408 super().__init__(handle_signals=handle_signals, **kwargs) 

409 self._app = app 

410 

411 @property 

412 def app(self) -> Application: 

413 return self._app 

414 

415 async def shutdown(self) -> None: 

416 await self._app.shutdown() 

417 

418 async def _make_server(self) -> Server: 

419 self._app.on_startup.freeze() 

420 await self._app.startup() 

421 self._app.freeze() 

422 

423 return Server( 

424 self._app._handle, # type: ignore[arg-type] 

425 request_factory=self._make_request, 

426 **self._kwargs, 

427 ) 

428 

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 ) 

448 

449 async def _cleanup_server(self) -> None: 

450 await self._app.cleanup()