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
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
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.
7from __future__ import annotations
9import asyncio
10import errno
11import socket
12import ssl
13import weakref
14from asyncio.trsock import TransportSocket
15from logging import getLogger
16from typing import Any, Callable
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
24_HAS_IPv6 = hasattr(socket, 'AF_INET6')
25_logger = getLogger('aiofastnet')
28def _is_asyncio_loop(loop: asyncio.AbstractEventLoop) -> bool:
29 return type(loop).__module__.startswith("asyncio.")
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')
37 if value is not None and value <= 0:
38 raise ValueError(f"{name} should be a positive number, got {value}")
40 if value is None:
41 return SSL_TIMEOUT_DEFAULTS[name]
43 return value
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')
51 if value is not None and value < 16384:
52 raise ValueError(f"{name} should be a positive number >= 16384, got {value}")
54 if value is None:
55 return SSL_BIO_SIZE_DEFAULTS[name]
57 return value
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)
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
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)
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
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 )
107 ssl_protocol_factory = ssl_transport.get_tls_protocol
109 create_connection = _get_original_loop_method(loop, "create_connection")
110 await create_connection(ssl_protocol_factory, None, None, sock=sock)
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)
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)
125 transport = wrapped_protocol._wrapped_transport
126 protocol = wrapped_protocol._protocol
127 wrapped_protocol._wrapped_transport = None
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)
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
180 return transport, protocol
183def _check_non_ssl_socket(sock):
184 if isinstance(sock, ssl.SSLSocket):
185 raise TypeError("Socket cannot be of type SSLSocket")
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)
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
207 if proto not in {0, socket.IPPROTO_TCP, socket.IPPROTO_UDP} or \
208 host is None:
209 return None
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
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
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]
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
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
256 # "host" is not an IP address.
257 return None
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
281 def __repr__(self):
282 return f'<{self.__class__.__name__} sockets={self.sockets!r}>'
284 def _attach(self, transport):
285 assert self._sockets is not None
286 self._clients.add(transport)
288 def _detach(self, transport):
289 self._clients.discard(transport)
290 if len(self._clients) == 0 and self._sockets is None:
291 self._wakeup()
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)
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)
310 def _start_serving_one_listener(self, listening_sock):
311 self._loop.add_reader(listening_sock.fileno(), self._accept_connection, listening_sock)
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))
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
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)
383 return
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
391 self._attach(transport)
393 def get_loop(self):
394 return self._loop
396 def is_serving(self):
397 return self._serving
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)
405 def close(self):
406 sockets = self._sockets
407 if sockets is None:
408 return
409 self._sockets = None
411 for sock in sockets:
412 self._loop.remove_reader(sock.fileno())
413 sock.close()
415 self._serving = False
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
422 if len(self._clients) == 0:
423 self._wakeup()
425 def close_clients(self):
426 for transport in self._clients.copy():
427 transport.close()
429 def abort_clients(self):
430 for transport in self._clients.copy():
431 transport.abort()
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)
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')
446 self._start_serving()
447 self._serving_forever_fut = self._loop.create_future()
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
460 async def wait_closed(self):
461 """Wait until server is closed and all connections are dropped.
463 - If the server is not closed, wait.
464 - If it is closed, but there are still active connections, wait.
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.
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
488def _stop_serving(loop, sock):
489 loop.remove_reader(sock.fileno())
490 sock.close()
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.')
504def _check_nonblocking_socket(py_sock):
505 if py_sock.getblocking():
506 raise ValueError("the socket must be non-blocking")