Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/aiofastnet/api_utils.py: 4%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

276 statements  

1# Portions of this file are derived from CPython's asyncio sources 

2# (notably asyncio.base_events and asyncio.selector_events). 

3# Copyright (c) Python Software Foundation. 

4# Licensed under the Python Software Foundation License Version 2. 

5# See LICENSES/PSF-2.0.txt and THIRD_PARTY_NOTICES for details. 

6 

7from __future__ import annotations 

8 

9import asyncio 

10import errno 

11import socket 

12import ssl 

13import weakref 

14from asyncio.trsock import TransportSocket 

15from logging import getLogger 

16from typing import Any, Callable 

17 

18from . import constants, openssl_compat 

19from .constants import SSL_BIO_SIZE_DEFAULTS, SSL_TIMEOUT_DEFAULTS 

20from .ssl_transport import SSLTransport_Socket, SSLTransport_Transport 

21from .transport import SelectorSocketTransport, aiofn_is_buffered_protocol 

22from .wrapped_transport import _get_original_loop_method, _should_fallback_to_asyncio, _WrappedBufferedProtocol, _WrappedProtocol 

23 

24_HAS_IPv6 = hasattr(socket, 'AF_INET6') 

25_logger = getLogger('aiofastnet') 

26 

27 

28def _is_asyncio_loop(loop: asyncio.AbstractEventLoop) -> bool: 

29 return type(loop).__module__.startswith("asyncio.") 

30 

31 

32def _validate_ssl_timeout(name: str, value: float | None, ssl_or_sslcontext: bool | ssl.SSLContext | None) -> float: 

33 if value is not None and not ssl_or_sslcontext: 

34 raise ValueError( 

35 f'{name} is only meaningful with ssl') 

36 

37 if value is not None and value <= 0: 

38 raise ValueError(f"{name} should be a positive number, got {value}") 

39 

40 if value is None: 

41 return SSL_TIMEOUT_DEFAULTS[name] 

42 

43 return value 

44 

45 

46def _validate_bio_size(name: str, value: int | None, ssl_or_sslcontext: bool | ssl.SSLContext | None) -> int: 

47 if value is not None and not ssl_or_sslcontext: 

48 raise ValueError( 

49 f'{name} is only meaningful with ssl') 

50 

51 if value is not None and value < 16384: 

52 raise ValueError(f"{name} should be a positive number >= 16384, got {value}") 

53 

54 if value is None: 

55 return SSL_BIO_SIZE_DEFAULTS[name] 

56 

57 return value 

58 

59 

60def _ssl_needs_fallback_engine(sslcontext: ssl.SSLContext) -> bool: 

61 return openssl_compat.OPENSSL_DYN_LIBS is None or getattr(sslcontext, "_aiofastnet_force_fallback_ssl", False) 

62 

63 

64async def _wait_and_close_transport_on_exc(waiter: asyncio.Future[Any], transport: Any) -> Any: 

65 try: 

66 return await waiter 

67 except: 

68 transport.close() 

69 raise 

70 

71 

72async def _create_connection_transport( 

73 loop: asyncio.AbstractEventLoop, 

74 sock: socket.socket, 

75 protocol_factory: Callable[[], asyncio.BaseProtocol], 

76 ssl: bool | ssl.SSLContext | None, 

77 server_hostname: str | None=None, 

78 server_side: bool=False, 

79 ssl_handshake_timeout: float | None=None, 

80 ssl_shutdown_timeout: float | None=None, 

81 ssl_incoming_bio_size: int | None=None, 

82 ssl_outgoing_bio_size: int | None=None, 

83 server=None 

84) -> tuple[asyncio.Transport, asyncio.BaseProtocol]: 

85 sock.setblocking(False) 

86 

87 # The following big nested if-else should set transport, protocol, and 

88 # optionally waiter variables 

89 waiter = None 

90 if _should_fallback_to_asyncio(loop): 

91 if ssl: 

92 protocol = protocol_factory() 

93 waiter = loop.create_future() if server is None else None 

94 sslcontext = None if isinstance(ssl, bool) else ssl 

95 

96 ssl_transport = SSLTransport_Transport( 

97 loop, protocol, sslcontext, 

98 server_side, 

99 ssl_handshake_timeout, 

100 ssl_shutdown_timeout, 

101 ssl_incoming_bio_size, 

102 ssl_outgoing_bio_size, 

103 waiter=waiter, 

104 server_hostname=server_hostname 

105 ) 

106 

107 ssl_protocol_factory = ssl_transport.get_tls_protocol 

108 

109 create_connection = _get_original_loop_method(loop, "create_connection") 

110 await create_connection(ssl_protocol_factory, None, None, sock=sock) 

111 

112 transport = ssl_transport 

113 else: 

114 def wrapped_protocol_factory(): 

115 user_protocol = protocol_factory() 

116 if aiofn_is_buffered_protocol(user_protocol): 

117 return _WrappedBufferedProtocol(user_protocol) 

118 else: 

119 return _WrappedProtocol(user_protocol) 

120 

121 create_connection = _get_original_loop_method(loop, "create_connection") 

122 loop_transport, wrapped_protocol = await create_connection( 

123 wrapped_protocol_factory, None, None, sock=sock) 

124 

125 transport = wrapped_protocol._wrapped_transport 

126 protocol = wrapped_protocol._protocol 

127 wrapped_protocol._wrapped_transport = None 

128 

129 if server is not None: 

130 # asyncio Transport needs _server in order to detach itself on disconnect 

131 loop_transport._server = server 

132 # and the Server must attach it, not WrappedTransport 

133 transport = loop_transport 

134 else: 

135 protocol = protocol_factory() 

136 waiter = loop.create_future() if server is None else None 

137 if ssl: 

138 sslcontext = openssl_compat.create_transport_context(server_side, server_hostname) if isinstance(ssl, bool) else ssl 

139 if _ssl_needs_fallback_engine(sslcontext): 

140 transport = SSLTransport_Transport( 

141 loop, protocol, sslcontext, 

142 server_side, 

143 ssl_handshake_timeout, 

144 ssl_shutdown_timeout, 

145 ssl_incoming_bio_size, 

146 ssl_outgoing_bio_size, 

147 waiter=waiter, 

148 server_hostname=server_hostname, 

149 server=server 

150 ) 

151 SelectorSocketTransport(loop, sock, transport.get_tls_protocol()) 

152 else: 

153 transport = SSLTransport_Socket( 

154 loop, protocol, sslcontext, 

155 server_side, 

156 ssl_handshake_timeout, 

157 ssl_shutdown_timeout, 

158 ssl_incoming_bio_size, 

159 ssl_outgoing_bio_size, 

160 sock, 

161 waiter=waiter, 

162 server_hostname=server_hostname, 

163 server=server 

164 ) 

165 else: 

166 transport = SelectorSocketTransport(loop, sock, protocol, 

167 waiter=waiter, server=server) 

168 

169 if waiter is not None: 

170 try: 

171 await waiter 

172 except: 

173 transport.close() 

174 # gh-109534: When an exception is raised by the SSLProtocol object the 

175 # exception set in this future can keep the protocol object alive and 

176 # cause a reference cycle. 

177 waiter = None 

178 raise 

179 

180 return transport, protocol 

181 

182 

183def _check_non_ssl_socket(sock): 

184 if isinstance(sock, ssl.SSLSocket): 

185 raise TypeError("Socket cannot be of type SSLSocket") 

186 

187 

188async def _ensure_resolved(address, *, 

189 family=0, type=socket.SOCK_STREAM, 

190 proto=0, flags=0, loop): 

191 host, port = address[:2] 

192 info = _ipaddr_info(host, port, family, type, proto, *address[2:]) 

193 if info is not None: 

194 # "host" is already a resolved IP. 

195 return [info] 

196 else: 

197 return await loop.getaddrinfo(host, port, family=family, type=type, 

198 proto=proto, flags=flags) 

199 

200 

201def _ipaddr_info(host, port, family, type, proto, flowinfo=0, scopeid=0): 

202 # Try to skip getaddrinfo if "host" is already an IP. Users might have 

203 # handled name resolution in their own code and pass in resolved IPs. 

204 if not hasattr(socket, 'inet_pton'): 

205 return 

206 

207 if proto not in {0, socket.IPPROTO_TCP, socket.IPPROTO_UDP} or \ 

208 host is None: 

209 return None 

210 

211 if type == socket.SOCK_STREAM: 

212 proto = socket.IPPROTO_TCP 

213 elif type == socket.SOCK_DGRAM: 

214 proto = socket.IPPROTO_UDP 

215 else: 

216 return None 

217 

218 if port is None: 

219 port = 0 

220 elif isinstance(port, bytes) and port == b'': 

221 port = 0 

222 elif isinstance(port, str) and port == '': 

223 port = 0 

224 else: 

225 # If port's a service name like "http", don't skip getaddrinfo. 

226 try: 

227 port = int(port) 

228 except (TypeError, ValueError): 

229 return None 

230 

231 if family == socket.AF_UNSPEC: 

232 afs = [socket.AF_INET] 

233 if _HAS_IPv6: 

234 afs.append(socket.AF_INET6) 

235 else: 

236 afs = [family] 

237 

238 if isinstance(host, bytes): 

239 host = host.decode('idna') 

240 if '%' in host: 

241 # Linux's inet_pton doesn't accept an IPv6 zone index after host, 

242 # like '::1%lo0'. 

243 return None 

244 

245 for af in afs: 

246 try: 

247 socket.inet_pton(af, host) 

248 # The host has already been resolved. 

249 if _HAS_IPv6 and af == socket.AF_INET6: 

250 return af, type, proto, '', (host, port, flowinfo, scopeid) 

251 else: 

252 return af, type, proto, '', (host, port) 

253 except OSError: 

254 pass 

255 

256 # "host" is not an IP address. 

257 return None 

258 

259 

260class Server(asyncio.AbstractServer): 

261 def __init__(self, loop, sockets, protocol_factory, ssl_context, backlog, 

262 ssl_handshake_timeout, ssl_shutdown_timeout, 

263 ssl_incoming_bio_size, ssl_outgoing_bio_size 

264 ): 

265 self._loop = loop 

266 self._sockets = sockets 

267 # Weak references so we don't break Transport's ability to 

268 # detect abandoned transports 

269 self._clients = weakref.WeakSet() 

270 self._waiters = [] 

271 self._protocol_factory = protocol_factory 

272 self._backlog = backlog 

273 self._ssl_context = ssl_context 

274 self._ssl_handshake_timeout = ssl_handshake_timeout 

275 self._ssl_shutdown_timeout = ssl_shutdown_timeout 

276 self._ssl_incoming_bio_size = ssl_incoming_bio_size 

277 self._ssl_outgoing_bio_size = ssl_outgoing_bio_size 

278 self._serving = False 

279 self._serving_forever_fut = None 

280 

281 def __repr__(self): 

282 return f'<{self.__class__.__name__} sockets={self.sockets!r}>' 

283 

284 def _attach(self, transport): 

285 assert self._sockets is not None 

286 self._clients.add(transport) 

287 

288 def _detach(self, transport): 

289 self._clients.discard(transport) 

290 if len(self._clients) == 0 and self._sockets is None: 

291 self._wakeup() 

292 

293 def _wakeup(self): 

294 waiters = self._waiters 

295 if waiters is None: 

296 return 

297 self._waiters = None 

298 for waiter in waiters: 

299 if not waiter.done(): 

300 waiter.set_result(None) 

301 

302 def _start_serving(self): 

303 if self._serving: 

304 return 

305 self._serving = True 

306 for sock in self._sockets: 

307 sock.listen(self._backlog) 

308 self._start_serving_one_listener(sock) 

309 

310 def _start_serving_one_listener(self, listening_sock): 

311 self._loop.add_reader(listening_sock.fileno(), self._accept_connection, listening_sock) 

312 

313 def _accept_connection(self, listening_sock): 

314 # This method is only called once for each event loop tick where the 

315 # listening socket has triggered an EVENT_READ. There may be multiple 

316 # connections waiting for an .accept() so it is called in a loop. 

317 # See https://bugs.python.org/issue27906 for more details. 

318 for _ in range(self._backlog + 1): 

319 try: 

320 conn, addr = listening_sock.accept() 

321 if self._loop.get_debug(): 

322 _logger.debug("%r got a new connection from %r: %r", self, addr, conn) 

323 conn.setblocking(False) 

324 except ConnectionAbortedError: 

325 # Discard connections that were aborted before accept(). 

326 continue 

327 except (BlockingIOError, InterruptedError): 

328 # Early exit because of a signal or 

329 # the socket accept buffer is empty. 

330 return 

331 except OSError as exc: 

332 # There's nowhere to send the error, so just log it. 

333 if exc.errno in (errno.EMFILE, errno.ENFILE, errno.ENOBUFS, errno.ENOMEM): 

334 # Some platforms (e.g. Linux keep reporting the FD as 

335 # ready, so we remove the read handler temporarily. 

336 # We'll try again in a while. 

337 self._loop.call_exception_handler( 

338 { 

339 "message": "socket.accept() out of system resource", 

340 "exception": exc, 

341 "socket": TransportSocket(listening_sock), 

342 } 

343 ) 

344 listening_sock.remove_reader(listening_sock.fileno()) 

345 self._loop.call_later(constants.ACCEPT_RETRY_DELAY, self._start_serving_one_listener, listening_sock) 

346 else: 

347 raise # The event loop will catch, log and ignore it. 

348 else: 

349 asyncio.create_task(self._accept_connection2(conn)) 

350 

351 async def _accept_connection2(self, sock): 

352 # By the time _accept_connection2 is called, server can be already closed 

353 # In such case we just close socket and return 

354 if self._sockets is None: 

355 sock.close() 

356 return 

357 

358 try: 

359 transport, _ = await _create_connection_transport( 

360 self._loop, 

361 sock, 

362 self._protocol_factory, 

363 self._ssl_context, 

364 server_hostname=None, 

365 server_side=True, 

366 ssl_handshake_timeout=self._ssl_handshake_timeout, 

367 ssl_shutdown_timeout=self._ssl_shutdown_timeout, 

368 ssl_incoming_bio_size=self._ssl_incoming_bio_size, 

369 ssl_outgoing_bio_size=self._ssl_outgoing_bio_size, 

370 server=self, 

371 ) 

372 except (SystemExit, KeyboardInterrupt): 

373 raise 

374 except BaseException as exc: 

375 sock.close() 

376 if self._loop.get_debug(): 

377 context = { 

378 "message": "Error on transport creation for incoming connection", 

379 "exception": exc, 

380 } 

381 self._loop.call_exception_handler(context) 

382 

383 return 

384 

385 # After await _create_connection_transport the server can be already closed 

386 # Abort transport and return then. 

387 if self._sockets is None: 

388 transport.abort() 

389 return 

390 

391 self._attach(transport) 

392 

393 def get_loop(self): 

394 return self._loop 

395 

396 def is_serving(self): 

397 return self._serving 

398 

399 @property 

400 def sockets(self): 

401 if self._sockets is None: 

402 return () 

403 return tuple(asyncio.trsock.TransportSocket(s) for s in self._sockets) 

404 

405 def close(self): 

406 sockets = self._sockets 

407 if sockets is None: 

408 return 

409 self._sockets = None 

410 

411 for sock in sockets: 

412 self._loop.remove_reader(sock.fileno()) 

413 sock.close() 

414 

415 self._serving = False 

416 

417 if (self._serving_forever_fut is not None and 

418 not self._serving_forever_fut.done()): 

419 self._serving_forever_fut.cancel() 

420 self._serving_forever_fut = None 

421 

422 if len(self._clients) == 0: 

423 self._wakeup() 

424 

425 def close_clients(self): 

426 for transport in self._clients.copy(): 

427 transport.close() 

428 

429 def abort_clients(self): 

430 for transport in self._clients.copy(): 

431 transport.abort() 

432 

433 async def start_serving(self): 

434 self._start_serving() 

435 # Skip one loop iteration so that all 'loop.add_reader' 

436 # go through. 

437 await asyncio.sleep(0) 

438 

439 async def serve_forever(self): 

440 if self._serving_forever_fut is not None: 

441 raise RuntimeError( 

442 f'server {self!r} is already being awaited on serve_forever()') 

443 if self._sockets is None: 

444 raise RuntimeError(f'server {self!r} is closed') 

445 

446 self._start_serving() 

447 self._serving_forever_fut = self._loop.create_future() 

448 

449 try: 

450 await self._serving_forever_fut 

451 except asyncio.CancelledError: 

452 try: 

453 self.close() 

454 await self.wait_closed() 

455 finally: 

456 raise 

457 finally: 

458 self._serving_forever_fut = None 

459 

460 async def wait_closed(self): 

461 """Wait until server is closed and all connections are dropped. 

462 

463 - If the server is not closed, wait. 

464 - If it is closed, but there are still active connections, wait. 

465 

466 Anyone waiting here will be unblocked once both conditions 

467 (server is closed and all connections have been dropped) 

468 have become true, in either order. 

469 

470 Historical note: In 3.11 and before, this was broken, returning 

471 immediately if the server was already closed, even if there 

472 were still active connections. An attempted fix in 3.12.0 was 

473 still broken, returning immediately if the server was still 

474 open and there were no active connections. Hopefully in 3.12.1 

475 we have it right. 

476 """ 

477 # Waiters are unblocked by self._wakeup(), which is called 

478 # from two places: self.close() and self._detach(), but only 

479 # when both conditions have become true. To signal that this 

480 # has happened, self._wakeup() sets self._waiters to None. 

481 if self._waiters is None: 

482 return 

483 waiter = self._loop.create_future() 

484 self._waiters.append(waiter) 

485 await waiter 

486 

487 

488def _stop_serving(loop, sock): 

489 loop.remove_reader(sock.fileno()) 

490 sock.close() 

491 

492 

493def _set_reuseport(sock): 

494 if not hasattr(socket, 'SO_REUSEPORT'): 

495 raise ValueError('reuse_port not supported by socket module') 

496 else: 

497 try: 

498 sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1) 

499 except OSError: 

500 raise ValueError('reuse_port not supported by socket module, ' 

501 'SO_REUSEPORT defined but not implemented.') 

502 

503 

504def _check_nonblocking_socket(py_sock): 

505 if py_sock.getblocking(): 

506 raise ValueError("the socket must be non-blocking")