Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
connector.py1855 linesDownload Raw Back to aiohttp
1import asyncio
2import functools
3import random
4import socket
5import sys
6import traceback
7import warnings
8from collections import OrderedDict, defaultdict, deque
9from contextlib import suppress
10from http import HTTPStatus
11from itertools import chain, cycle, islice
12from time import monotonic
13from types import TracebackType
14from typing import (
15    TYPE_CHECKING,
16    Any,
17    Awaitable,
18    Callable,
19    DefaultDict,
20    Deque,
21    Dict,
22    Iterator,
23    List,
24    Literal,
25    Optional,
26    Sequence,
27    Set,
28    Tuple,
29    Type,
30    Union,
31    cast,
32)
33
34import aiohappyeyeballs
35from aiohappyeyeballs import AddrInfoType, SocketFactoryType
36
37from . import hdrs, helpers
38from .abc import AbstractResolver, ResolveResult
39from .client_exceptions import (
40    ClientConnectionError,
41    ClientConnectorCertificateError,
42    ClientConnectorDNSError,
43    ClientConnectorError,
44    ClientConnectorSSLError,
45    ClientHttpProxyError,
46    ClientProxyConnectionError,
47    ServerFingerprintMismatch,
48    UnixClientConnectorError,
49    cert_errors,
50    ssl_errors,
51)
52from .client_proto import ResponseHandler
53from .client_reqrep import ClientRequest, Fingerprint, _merge_ssl_params
54from .helpers import (
55    _SENTINEL,
56    ceil_timeout,
57    is_ip_address,
58    noop,
59    sentinel,
60    set_exception,
61    set_result,
62)
63from .log import client_logger
64from .resolver import DefaultResolver
65
66if sys.version_info >= (3, 12):
67    from collections.abc import Buffer
68else:
69    Buffer = Union[bytes, bytearray, "memoryview[int]", "memoryview[bytes]"]
70
71if TYPE_CHECKING:
72    import ssl
73
74    SSLContext = ssl.SSLContext
75else:
76    try:
77        import ssl
78
79        SSLContext = ssl.SSLContext
80    except ImportError:  # pragma: no cover
81        ssl = None  # type: ignore[assignment]
82        SSLContext = object  # type: ignore[misc,assignment]
83
84EMPTY_SCHEMA_SET = frozenset({""})
85HTTP_SCHEMA_SET = frozenset({"http", "https"})
86WS_SCHEMA_SET = frozenset({"ws", "wss"})
87
88HTTP_AND_EMPTY_SCHEMA_SET = HTTP_SCHEMA_SET | EMPTY_SCHEMA_SET
89HIGH_LEVEL_SCHEMA_SET = HTTP_AND_EMPTY_SCHEMA_SET | WS_SCHEMA_SET
90
91NEEDS_CLEANUP_CLOSED = (3, 13, 0) <= sys.version_info < (
92    3,
93    13,
94    1,
95) or sys.version_info < (3, 12, 8)
96# Cleanup closed is no longer needed after https://github.com/python/cpython/pull/118960
97# which first appeared in Python 3.12.8 and 3.13.1
98
99
100__all__ = (
101    "BaseConnector",
102    "TCPConnector",
103    "UnixConnector",
104    "NamedPipeConnector",
105    "AddrInfoType",
106    "SocketFactoryType",
107)
108
109
110if TYPE_CHECKING:
111    from .client import ClientTimeout
112    from .client_reqrep import ConnectionKey
113    from .tracing import Trace
114
115
116class _DeprecationWaiter:
117    __slots__ = ("_awaitable", "_awaited")
118
119    def __init__(self, awaitable: Awaitable[Any]) -> None:
120        self._awaitable = awaitable
121        self._awaited = False
122
123    def __await__(self) -> Any:
124        self._awaited = True
125        return self._awaitable.__await__()
126
127    def __del__(self) -> None:
128        if not self._awaited:
129            warnings.warn(
130                "Connector.close() is a coroutine, "
131                "please use await connector.close()",
132                DeprecationWarning,
133            )
134
135
136async def _wait_for_close(waiters: List[Awaitable[object]]) -> None:
137    """Wait for all waiters to finish closing."""
138    results = await asyncio.gather(*waiters, return_exceptions=True)
139    for res in results:
140        if isinstance(res, Exception):
141            client_logger.debug("Error while closing connector: %r", res)
142
143
144class Connection:
145
146    _source_traceback = None
147
148    def __init__(
149        self,
150        connector: "BaseConnector",
151        key: "ConnectionKey",
152        protocol: ResponseHandler,
153        loop: asyncio.AbstractEventLoop,
154    ) -> None:
155        self._key = key
156        self._connector = connector
157        self._loop = loop
158        self._protocol: Optional[ResponseHandler] = protocol
159        self._callbacks: List[Callable[[], None]] = []
160
161        if loop.get_debug():
162            self._source_traceback = traceback.extract_stack(sys._getframe(1))
163
164    def __repr__(self) -> str:
165        return f"Connection<{self._key}>"
166
167    def __del__(self, _warnings: Any = warnings) -> None:
168        if self._protocol is not None:
169            kwargs = {"source": self}
170            _warnings.warn(f"Unclosed connection {self!r}", ResourceWarning, **kwargs)
171            if self._loop.is_closed():
172                return
173
174            self._connector._release(self._key, self._protocol, should_close=True)
175
176            context = {"client_connection": self, "message": "Unclosed connection"}
177            if self._source_traceback is not None:
178                context["source_traceback"] = self._source_traceback
179            self._loop.call_exception_handler(context)
180
181    def __bool__(self) -> Literal[True]:
182        """Force subclasses to not be falsy, to make checks simpler."""
183        return True
184
185    @property
186    def loop(self) -> asyncio.AbstractEventLoop:
187        warnings.warn(
188            "connector.loop property is deprecated", DeprecationWarning, stacklevel=2
189        )
190        return self._loop
191
192    @property
193    def transport(self) -> Optional[asyncio.Transport]:
194        if self._protocol is None:
195            return None
196        return self._protocol.transport
197
198    @property
199    def protocol(self) -> Optional[ResponseHandler]:
200        return self._protocol
201
202    def add_callback(self, callback: Callable[[], None]) -> None:
203        if callback is not None:
204            self._callbacks.append(callback)
205
206    def _notify_release(self) -> None:
207        callbacks, self._callbacks = self._callbacks[:], []
208
209        for cb in callbacks:
210            with suppress(Exception):
211                cb()
212
213    def close(self) -> None:
214        self._notify_release()
215
216        if self._protocol is not None:
217            self._connector._release(self._key, self._protocol, should_close=True)
218            self._protocol = None
219
220    def release(self) -> None:
221        self._notify_release()
222
223        if self._protocol is not None:
224            self._connector._release(self._key, self._protocol)
225            self._protocol = None
226
227    @property
228    def closed(self) -> bool:
229        return self._protocol is None or not self._protocol.is_connected()
230
231
232class _ConnectTunnelConnection(Connection):
233    """Special connection wrapper for CONNECT tunnels that must never be pooled.
234
235    This connection wraps the proxy connection that will be upgraded with TLS.
236    It must never be released to the pool because:
237    1. Its 'closed' future will never complete, causing session.close() to hang
238    2. It represents an intermediate state, not a reusable connection
239    3. The real connection (with TLS) will be created separately
240    """
241
242    def release(self) -> None:
243        """Do nothing - don't pool or close the connection.
244
245        These connections are an intermediate state during the CONNECT tunnel
246        setup and will be cleaned up naturally after the TLS upgrade. If they
247        were to be pooled, they would never be properly closed, causing
248        session.close() to wait forever for their 'closed' future.
249        """
250
251
252class _TransportPlaceholder:
253    """placeholder for BaseConnector.connect function"""
254
255    __slots__ = ("closed", "transport")
256
257    def __init__(self, closed_future: asyncio.Future[Optional[Exception]]) -> None:
258        """Initialize a placeholder for a transport."""
259        self.closed = closed_future
260        self.transport = None
261
262    def close(self) -> None:
263        """Close the placeholder."""
264
265    def abort(self) -> None:
266        """Abort the placeholder (does nothing)."""
267
268
269class BaseConnector:
270    """Base connector class.
271
272    keepalive_timeout - (optional) Keep-alive timeout.
273    force_close - Set to True to force close and do reconnect
274        after each request (and between redirects).
275    limit - The total number of simultaneous connections.
276    limit_per_host - Number of simultaneous connections to one host.
277    enable_cleanup_closed - Enables clean-up closed ssl transports.
278                            Disabled by default.
279    timeout_ceil_threshold - Trigger ceiling of timeout values when
280                             it's above timeout_ceil_threshold.
281    loop - Optional event loop.
282    """
283
284    _closed = True  # prevent AttributeError in __del__ if ctor was failed
285    _source_traceback = None
286
287    # abort transport after 2 seconds (cleanup broken connections)
288    _cleanup_closed_period = 2.0
289
290    allowed_protocol_schema_set = HIGH_LEVEL_SCHEMA_SET
291
292    def __init__(
293        self,
294        *,
295        keepalive_timeout: Union[object, None, float] = sentinel,
296        force_close: bool = False,
297        limit: int = 100,
298        limit_per_host: int = 0,
299        enable_cleanup_closed: bool = False,
300        loop: Optional[asyncio.AbstractEventLoop] = None,
301        timeout_ceil_threshold: float = 5,
302    ) -> None:
303
304        if force_close:
305            if keepalive_timeout is not None and keepalive_timeout is not sentinel:
306                raise ValueError(
307                    "keepalive_timeout cannot be set if force_close is True"
308                )
309        else:
310            if keepalive_timeout is sentinel:
311                keepalive_timeout = 15.0
312
313        loop = loop or asyncio.get_running_loop()
314        self._timeout_ceil_threshold = timeout_ceil_threshold
315
316        self._closed = False
317        if loop.get_debug():
318            self._source_traceback = traceback.extract_stack(sys._getframe(1))
319
320        # Connection pool of reusable connections.
321        # We use a deque to store connections because it has O(1) popleft()
322        # and O(1) append() operations to implement a FIFO queue.
323        self._conns: DefaultDict[
324            ConnectionKey, Deque[Tuple[ResponseHandler, float]]
325        ] = defaultdict(deque)
326        self._limit = limit
327        self._limit_per_host = limit_per_host
328        self._acquired: Set[ResponseHandler] = set()
329        self._acquired_per_host: DefaultDict[ConnectionKey, Set[ResponseHandler]] = (
330            defaultdict(set)
331        )
332        self._keepalive_timeout = cast(float, keepalive_timeout)
333        self._force_close = force_close
334
335        # {host_key: FIFO list of waiters}
336        # The FIFO is implemented with an OrderedDict with None keys because
337        # python does not have an ordered set.
338        self._waiters: DefaultDict[
339            ConnectionKey, OrderedDict[asyncio.Future[None], None]
340        ] = defaultdict(OrderedDict)
341
342        self._loop = loop
343        self._factory = functools.partial(ResponseHandler, loop=loop)
344
345        # start keep-alive connection cleanup task
346        self._cleanup_handle: Optional[asyncio.TimerHandle] = None
347
348        # start cleanup closed transports task
349        self._cleanup_closed_handle: Optional[asyncio.TimerHandle] = None
350
351        if enable_cleanup_closed and not NEEDS_CLEANUP_CLOSED:
352            warnings.warn(
353                "enable_cleanup_closed ignored because "
354                "https://github.com/python/cpython/pull/118960 is fixed "
355                f"in Python version {sys.version_info}",
356                DeprecationWarning,
357                stacklevel=2,
358            )
359            enable_cleanup_closed = False
360
361        self._cleanup_closed_disabled = not enable_cleanup_closed
362        self._cleanup_closed_transports: List[Optional[asyncio.Transport]] = []
363        self._placeholder_future: asyncio.Future[Optional[Exception]] = (
364            loop.create_future()
365        )
366        self._placeholder_future.set_result(None)
367        self._cleanup_closed()
368
369    def __del__(self, _warnings: Any = warnings) -> None:
370        if self._closed:
371            return
372        if not self._conns:
373            return
374
375        conns = [repr(c) for c in self._conns.values()]
376
377        self._close()
378
379        kwargs = {"source": self}
380        _warnings.warn(f"Unclosed connector {self!r}", ResourceWarning, **kwargs)
381        context = {
382            "connector": self,
383            "connections": conns,
384            "message": "Unclosed connector",
385        }
386        if self._source_traceback is not None:
387            context["source_traceback"] = self._source_traceback
388        self._loop.call_exception_handler(context)
389
390    def __enter__(self) -> "BaseConnector":
391        warnings.warn(
392            '"with Connector():" is deprecated, '
393            'use "async with Connector():" instead',
394            DeprecationWarning,
395        )
396        return self
397
398    def __exit__(self, *exc: Any) -> None:
399        self._close()
400
401    async def __aenter__(self) -> "BaseConnector":
402        return self
403
404    async def __aexit__(
405        self,
406        exc_type: Optional[Type[BaseException]] = None,
407        exc_value: Optional[BaseException] = None,
408        exc_traceback: Optional[TracebackType] = None,
409    ) -> None:
410        await self.close()
411
412    @property
413    def force_close(self) -> bool:
414        """Ultimately close connection on releasing if True."""
415        return self._force_close
416
417    @property
418    def limit(self) -> int:
419        """The total number for simultaneous connections.
420
421        If limit is 0 the connector has no limit.
422        The default limit size is 100.
423        """
424        return self._limit
425
426    @property
427    def limit_per_host(self) -> int:
428        """The limit for simultaneous connections to the same endpoint.
429
430        Endpoints are the same if they are have equal
431        (host, port, is_ssl) triple.
432        """
433        return self._limit_per_host
434
435    def _cleanup(self) -> None:
436        """Cleanup unused transports."""
437        if self._cleanup_handle:
438            self._cleanup_handle.cancel()
439            # _cleanup_handle should be unset, otherwise _release() will not
440            # recreate it ever!
441            self._cleanup_handle = None
442
443        now = monotonic()
444        timeout = self._keepalive_timeout
445
446        if self._conns:
447            connections = defaultdict(deque)
448            deadline = now - timeout
449            for key, conns in self._conns.items():
450                alive: Deque[Tuple[ResponseHandler, float]] = deque()
451                for proto, use_time in conns:
452                    if proto.is_connected() and use_time - deadline >= 0:
453                        alive.append((proto, use_time))
454                        continue
455                    transport = proto.transport
456                    proto.close()
457                    if not self._cleanup_closed_disabled and key.is_ssl:
458                        self._cleanup_closed_transports.append(transport)
459
460                if alive:
461                    connections[key] = alive
462
463            self._conns = connections
464
465        if self._conns:
466            self._cleanup_handle = helpers.weakref_handle(
467                self,
468                "_cleanup",
469                timeout,
470                self._loop,
471                timeout_ceil_threshold=self._timeout_ceil_threshold,
472            )
473
474    def _cleanup_closed(self) -> None:
475        """Double confirmation for transport close.
476
477        Some broken ssl servers may leave socket open without proper close.
478        """
479        if self._cleanup_closed_handle:
480            self._cleanup_closed_handle.cancel()
481
482        for transport in self._cleanup_closed_transports:
483            if transport is not None:
484                transport.abort()
485
486        self._cleanup_closed_transports = []
487
488        if not self._cleanup_closed_disabled:
489            self._cleanup_closed_handle = helpers.weakref_handle(
490                self,
491                "_cleanup_closed",
492                self._cleanup_closed_period,
493                self._loop,
494                timeout_ceil_threshold=self._timeout_ceil_threshold,
495            )
496
497    def close(self, *, abort_ssl: bool = False) -> Awaitable[None]:
498        """Close all opened transports.
499
500        :param abort_ssl: If True, SSL connections will be aborted immediately
501                         without performing the shutdown handshake. This provides
502                         faster cleanup at the cost of less graceful disconnection.
503        """
504        if not (waiters := self._close(abort_ssl=abort_ssl)):
505            # If there are no connections to close, we can return a noop
506            # awaitable to avoid scheduling a task on the event loop.
507            return _DeprecationWaiter(noop())
508        coro = _wait_for_close(waiters)
509        if sys.version_info >= (3, 12):
510            # Optimization for Python 3.12, try to close connections
511            # immediately to avoid having to schedule the task on the event loop.
512            task = asyncio.Task(coro, loop=self._loop, eager_start=True)
513        else:
514            task = self._loop.create_task(coro)
515        return _DeprecationWaiter(task)
516
517    def _close(self, *, abort_ssl: bool = False) -> List[Awaitable[object]]:
518        waiters: List[Awaitable[object]] = []
519
520        if self._closed:
521            return waiters
522
523        self._closed = True
524
525        try:
526            if self._loop.is_closed():
527                return waiters
528
529            # cancel cleanup task
530            if self._cleanup_handle:
531                self._cleanup_handle.cancel()
532
533            # cancel cleanup close task
534            if self._cleanup_closed_handle:
535                self._cleanup_closed_handle.cancel()
536
537            for data in self._conns.values():
538                for proto, _ in data:
539                    if (
540                        abort_ssl
541                        and proto.transport
542                        and proto.transport.get_extra_info("sslcontext") is not None
543                    ):
544                        proto.abort()
545                    else:
546                        proto.close()
547                    if closed := proto.closed:
548                        waiters.append(closed)
549
550            for proto in self._acquired:
551                if (
552                    abort_ssl
553                    and proto.transport
554                    and proto.transport.get_extra_info("sslcontext") is not None
555                ):
556                    proto.abort()
557                else:
558                    proto.close()
559                if closed := proto.closed:
560                    waiters.append(closed)
561
562            for transport in self._cleanup_closed_transports:
563                if transport is not None:
564                    transport.abort()
565
566            return waiters
567
568        finally:
569            self._conns.clear()
570            self._acquired.clear()
571            for keyed_waiters in self._waiters.values():
572                for keyed_waiter in keyed_waiters:
573                    keyed_waiter.cancel()
574            self._waiters.clear()
575            self._cleanup_handle = None
576            self._cleanup_closed_transports.clear()
577            self._cleanup_closed_handle = None
578
579    @property
580    def closed(self) -> bool:
581        """Is connector closed.
582
583        A readonly property.
584        """
585        return self._closed
586
587    def _available_connections(self, key: "ConnectionKey") -> int:
588        """
589        Return number of available connections.
590
591        The limit, limit_per_host and the connection key are taken into account.
592
593        If it returns less than 1 means that there are no connections
594        available.
595        """
596        # check total available connections
597        # If there are no limits, this will always return 1
598        total_remain = 1
599
600        if self._limit and (total_remain := self._limit - len(self._acquired)) <= 0:
601            return total_remain
602
603        # check limit per host
604        if host_remain := self._limit_per_host:
605            if acquired := self._acquired_per_host.get(key):
606                host_remain -= len(acquired)
607            if total_remain > host_remain:
608                return host_remain
609
610        return total_remain
611
612    def _update_proxy_auth_header_and_build_proxy_req(
613        self, req: ClientRequest
614    ) -> ClientRequest:
615        """Set Proxy-Authorization header for non-SSL proxy requests and builds the proxy request for SSL proxy requests."""
616        url = req.proxy
617        assert url is not None
618        headers: Dict[str, str] = {}
619        if req.proxy_headers is not None:
620            headers = req.proxy_headers  # type: ignore[assignment]
621        headers[hdrs.HOST] = req.headers[hdrs.HOST]
622        proxy_req = ClientRequest(
623            hdrs.METH_GET,
624            url,
625            headers=headers,
626            auth=req.proxy_auth,
627            loop=self._loop,
628            ssl=req.ssl,
629        )
630        auth = proxy_req.headers.pop(hdrs.AUTHORIZATION, None)
631        if auth is not None:
632            if not req.is_ssl():
633                req.headers[hdrs.PROXY_AUTHORIZATION] = auth
634            else:
635                proxy_req.headers[hdrs.PROXY_AUTHORIZATION] = auth
636        return proxy_req
637
638    async def connect(
639        self, req: ClientRequest, traces: List["Trace"], timeout: "ClientTimeout"
640    ) -> Connection:
641        """Get from pool or create new connection."""
642        key = req.connection_key
643        if (conn := await self._get(key, traces)) is not None:
644            # If we do not have to wait and we can get a connection from the pool
645            # we can avoid the timeout ceil logic and directly return the connection
646            if req.proxy:
647                self._update_proxy_auth_header_and_build_proxy_req(req)
648            return conn
649
650        async with ceil_timeout(timeout.connect, timeout.ceil_threshold):
651            if self._available_connections(key) <= 0:
652                await self._wait_for_available_connection(key, traces)
653                if (conn := await self._get(key, traces)) is not None:
654                    if req.proxy:
655                        self._update_proxy_auth_header_and_build_proxy_req(req)
656                    return conn
657
658            placeholder = cast(
659                ResponseHandler, _TransportPlaceholder(self._placeholder_future)
660            )
661            self._acquired.add(placeholder)
662            if self._limit_per_host:
663                self._acquired_per_host[key].add(placeholder)
664
665            try:
666                # Traces are done inside the try block to ensure that the
667                # that the placeholder is still cleaned up if an exception
668                # is raised.
669                if traces:
670                    for trace in traces:
671                        await trace.send_connection_create_start()
672                proto = await self._create_connection(req, traces, timeout)
673                if traces:
674                    for trace in traces:
675                        await trace.send_connection_create_end()
676            except BaseException:
677                self._release_acquired(key, placeholder)
678                raise
679            else:
680                if self._closed:
681                    proto.close()
682                    raise ClientConnectionError("Connector is closed.")
683
684        # The connection was successfully created, drop the placeholder
685        # and add the real connection to the acquired set. There should
686        # be no awaits after the proto is added to the acquired set
687        # to ensure that the connection is not left in the acquired set
688        # on cancellation.
689        self._acquired.remove(placeholder)
690        self._acquired.add(proto)
691        if self._limit_per_host:
692            acquired_per_host = self._acquired_per_host[key]
693            acquired_per_host.remove(placeholder)
694            acquired_per_host.add(proto)
695        return Connection(self, key, proto, self._loop)
696
697    async def _wait_for_available_connection(
698        self, key: "ConnectionKey", traces: List["Trace"]
699    ) -> None:
700        """Wait for an available connection slot."""
701        # We loop here because there is a race between
702        # the connection limit check and the connection
703        # being acquired. If the connection is acquired
704        # between the check and the await statement, we
705        # need to loop again to check if the connection
706        # slot is still available.
707        attempts = 0
708        while True:
709            fut: asyncio.Future[None] = self._loop.create_future()
710            keyed_waiters = self._waiters[key]
711            keyed_waiters[fut] = None
712            if attempts:
713                # If we have waited before, we need to move the waiter
714                # to the front of the queue as otherwise we might get
715                # starved and hit the timeout.
716                keyed_waiters.move_to_end(fut, last=False)
717
718            try:
719                # Traces happen in the try block to ensure that the
720                # the waiter is still cleaned up if an exception is raised.
721                if traces:
722                    for trace in traces:
723                        await trace.send_connection_queued_start()
724                await fut
725                if traces:
726                    for trace in traces:
727                        await trace.send_connection_queued_end()
728            finally:
729                # pop the waiter from the queue if its still
730                # there and not already removed by _release_waiter
731                keyed_waiters.pop(fut, None)
732                if not self._waiters.get(key, True):
733                    del self._waiters[key]
734
735            if self._available_connections(key) > 0:
736                break
737            attempts += 1
738
739    async def _get(
740        self, key: "ConnectionKey", traces: List["Trace"]
741    ) -> Optional[Connection]:
742        """Get next reusable connection for the key or None.
743
744        The connection will be marked as acquired.
745        """
746        if (conns := self._conns.get(key)) is None:
747            return None
748
749        t1 = monotonic()
750        while conns:
751            proto, t0 = conns.popleft()
752            # We will we reuse the connection if its connected and
753            # the keepalive timeout has not been exceeded
754            if proto.is_connected() and t1 - t0 <= self._keepalive_timeout:
755                if not conns:
756                    # The very last connection was reclaimed: drop the key
757                    del self._conns[key]
758                self._acquired.add(proto)
759                if self._limit_per_host:
760                    self._acquired_per_host[key].add(proto)
761                if traces:
762                    for trace in traces:
763                        try:
764                            await trace.send_connection_reuseconn()
765                        except BaseException:
766                            self._release_acquired(key, proto)
767                            raise
768                return Connection(self, key, proto, self._loop)
769
770            # Connection cannot be reused, close it
771            transport = proto.transport
772            proto.close()
773            # only for SSL transports
774            if not self._cleanup_closed_disabled and key.is_ssl:
775                self._cleanup_closed_transports.append(transport)
776
777        # No more connections: drop the key
778        del self._conns[key]
779        return None
780
781    def _release_waiter(self) -> None:
782        """
783        Iterates over all waiters until one to be released is found.
784
785        The one to be released is not finished and
786        belongs to a host that has available connections.
787        """
788        if not self._waiters:
789            return
790
791        # Having the dict keys ordered this avoids to iterate
792        # at the same order at each call.
793        queues = list(self._waiters)
794        random.shuffle(queues)
795
796        for key in queues:
797            if self._available_connections(key) < 1:
798                continue
799
800            waiters = self._waiters[key]
801            while waiters:
802                waiter, _ = waiters.popitem(last=False)
803                if not waiter.done():
804                    waiter.set_result(None)
805                    return
806
807    def _release_acquired(self, key: "ConnectionKey", proto: ResponseHandler) -> None:
808        """Release acquired connection."""
809        if self._closed:
810            # acquired connection is already released on connector closing
811            return
812
813        self._acquired.discard(proto)
814        if self._limit_per_host and (conns := self._acquired_per_host.get(key)):
815            conns.discard(proto)
816            if not conns:
817                del self._acquired_per_host[key]
818        self._release_waiter()
819
820    def _release(
821        self,
822        key: "ConnectionKey",
823        protocol: ResponseHandler,
824        *,
825        should_close: bool = False,
826    ) -> None:
827        if self._closed:
828            # acquired connection is already released on connector closing
829            return
830
831        self._release_acquired(key, protocol)
832
833        if self._force_close or should_close or protocol.should_close:
834            transport = protocol.transport
835            protocol.close()
836
837            if key.is_ssl and not self._cleanup_closed_disabled:
838                self._cleanup_closed_transports.append(transport)
839            return
840
841        self._conns[key].append((protocol, monotonic()))
842
843        if self._cleanup_handle is None:
844            self._cleanup_handle = helpers.weakref_handle(
845                self,
846                "_cleanup",
847                self._keepalive_timeout,
848                self._loop,
849                timeout_ceil_threshold=self._timeout_ceil_threshold,
850            )
851
852    async def _create_connection(
853        self, req: ClientRequest, traces: List["Trace"], timeout: "ClientTimeout"
854    ) -> ResponseHandler:
855        raise NotImplementedError()
856
857
858class _DNSCacheTable:
859    def __init__(self, ttl: Optional[float] = None, max_size: int = 1000) -> None:
860        self._addrs_rr: OrderedDict[
861            Tuple[str, int], Tuple[Iterator[ResolveResult], int]
862        ] = OrderedDict()
863        self._timestamps: Dict[Tuple[str, int], float] = {}
864        self._ttl = ttl
865        self._max_size = max_size
866
867    def __contains__(self, host: object) -> bool:
868        return host in self._addrs_rr
869
870    def add(self, key: Tuple[str, int], addrs: List[ResolveResult]) -> None:
871        if key in self._addrs_rr:
872            self._addrs_rr.move_to_end(key)
873
874        self._addrs_rr[key] = (cycle(addrs), len(addrs))
875
876        if self._ttl is not None:
877            self._timestamps[key] = monotonic()
878
879        if len(self._addrs_rr) > self._max_size:
880            oldest_key, _ = self._addrs_rr.popitem(last=False)
881            self._timestamps.pop(oldest_key, None)
882
883    def remove(self, key: Tuple[str, int]) -> None:
884        self._addrs_rr.pop(key, None)
885        self._timestamps.pop(key, None)
886
887    def clear(self) -> None:
888        self._addrs_rr.clear()
889        self._timestamps.clear()
890
891    def next_addrs(self, key: Tuple[str, int]) -> List[ResolveResult]:
892        loop, length = self._addrs_rr[key]
893        addrs = list(islice(loop, length))
894        # Consume one more element to shift internal state of `cycle`
895        next(loop)
896        self._addrs_rr.move_to_end(key)
897        return addrs
898
899    def expired(self, key: Tuple[str, int]) -> bool:
900        if self._ttl is None:
901            return False
902
903        return self._timestamps[key] + self._ttl < monotonic()
904
905
906def _make_ssl_context(verified: bool) -> SSLContext:
907    """Create SSL context.
908
909    This method is not async-friendly and should be called from a thread
910    because it will load certificates from disk and do other blocking I/O.
911    """
912    if ssl is None:
913        # No ssl support
914        return None
915    if verified:
916        sslcontext = ssl.create_default_context()
917    else:
918        sslcontext = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
919        sslcontext.options |= ssl.OP_NO_SSLv2
920        sslcontext.options |= ssl.OP_NO_SSLv3
921        sslcontext.check_hostname = False
922        sslcontext.verify_mode = ssl.CERT_NONE
923        sslcontext.options |= ssl.OP_NO_COMPRESSION
924        sslcontext.set_default_verify_paths()
925    sslcontext.set_alpn_protocols(("http/1.1",))
926    return sslcontext
927
928
929# The default SSLContext objects are created at import time
930# since they do blocking I/O to load certificates from disk,
931# and imports should always be done before the event loop starts
932# or in a thread.
933_SSL_CONTEXT_VERIFIED = _make_ssl_context(True)
934_SSL_CONTEXT_UNVERIFIED = _make_ssl_context(False)
935
936
937class TCPConnector(BaseConnector):
938    """TCP connector.
939
940    verify_ssl - Set to True to check ssl certifications.
941    fingerprint - Pass the binary sha256
942        digest of the expected certificate in DER format to verify
943        that the certificate the server presents matches. See also
944        https://en.wikipedia.org/wiki/HTTP_Public_Key_Pinning
945    resolver - Enable DNS lookups and use this
946        resolver
947    use_dns_cache - Use memory cache for DNS lookups.
948    ttl_dns_cache - Max seconds having cached a DNS entry, None forever.
949    family - socket address family
950    local_addr - local tuple of (host, port) to bind socket to
951
952    keepalive_timeout - (optional) Keep-alive timeout.
953    force_close - Set to True to force close and do reconnect
954        after each request (and between redirects).
955    limit - The total number of simultaneous connections.
956    limit_per_host - Number of simultaneous connections to one host.
957    enable_cleanup_closed - Enables clean-up closed ssl transports.
958                            Disabled by default.
959    happy_eyeballs_delay - This is the “Connection Attempt Delay”
960                           as defined in RFC 8305. To disable
961                           the happy eyeballs algorithm, set to None.
962    interleave - “First Address Family Count” as defined in RFC 8305
963    loop - Optional event loop.
964    socket_factory - A SocketFactoryType function that, if supplied,
965                     will be used to create sockets given an
966                     AddrInfoType.
967    ssl_shutdown_timeout - DEPRECATED. Will be removed in aiohttp 4.0.
968                           Grace period for SSL shutdown handshake on TLS
969                           connections. Default is 0 seconds (immediate abort).
970                           This parameter allowed for a clean SSL shutdown by
971                           notifying the remote peer of connection closure,
972                           while avoiding excessive delays during connector cleanup.
973                           Note: Only takes effect on Python 3.11+.
974    """
975
976    allowed_protocol_schema_set = HIGH_LEVEL_SCHEMA_SET | frozenset({"tcp"})
977
978    def __init__(
979        self,
980        *,
981        verify_ssl: bool = True,
982        fingerprint: Optional[bytes] = None,
983        use_dns_cache: bool = True,
984        ttl_dns_cache: Optional[int] = 10,
985        dns_cache_max_size: int = 1000,
986        family: socket.AddressFamily = socket.AddressFamily.AF_UNSPEC,
987        ssl_context: Optional[SSLContext] = None,
988        ssl: Union[bool, Fingerprint, SSLContext] = True,
989        local_addr: Optional[Tuple[str, int]] = None,
990        resolver: Optional[AbstractResolver] = None,
991        keepalive_timeout: Union[None, float, object] = sentinel,
992        force_close: bool = False,
993        limit: int = 100,
994        limit_per_host: int = 0,
995        enable_cleanup_closed: bool = False,
996        loop: Optional[asyncio.AbstractEventLoop] = None,
997        timeout_ceil_threshold: float = 5,
998        happy_eyeballs_delay: Optional[float] = 0.25,
999        interleave: Optional[int] = None,
1000        socket_factory: Optional[SocketFactoryType] = None,
1001        ssl_shutdown_timeout: Union[_SENTINEL, None, float] = sentinel,
1002    ):
1003        super().__init__(
1004            keepalive_timeout=keepalive_timeout,
1005            force_close=force_close,
1006            limit=limit,
1007            limit_per_host=limit_per_host,
1008            enable_cleanup_closed=enable_cleanup_closed,
1009            loop=loop,
1010            timeout_ceil_threshold=timeout_ceil_threshold,
1011        )
1012
1013        self._ssl = _merge_ssl_params(ssl, verify_ssl, ssl_context, fingerprint)
1014
1015        self._resolver: AbstractResolver
1016        if resolver is None:
1017            self._resolver = DefaultResolver(loop=self._loop)
1018            self._resolver_owner = True
1019        else:
1020            self._resolver = resolver
1021            self._resolver_owner = False
1022
1023        self._use_dns_cache = use_dns_cache
1024        self._cached_hosts = _DNSCacheTable(
1025            ttl=ttl_dns_cache, max_size=dns_cache_max_size
1026        )
1027        self._throttle_dns_futures: Dict[
1028            Tuple[str, int], Set["asyncio.Future[None]"]
1029        ] = {}
1030        self._family = family
1031        self._local_addr_infos = aiohappyeyeballs.addr_to_addr_infos(local_addr)
1032        self._happy_eyeballs_delay = happy_eyeballs_delay
1033        self._interleave = interleave
1034        self._resolve_host_tasks: Set["asyncio.Task[List[ResolveResult]]"] = set()
1035        self._socket_factory = socket_factory
1036        self._ssl_shutdown_timeout: Optional[float]
1037        # Handle ssl_shutdown_timeout with warning for Python < 3.11
1038        if ssl_shutdown_timeout is sentinel:
1039            self._ssl_shutdown_timeout = 0
1040        else:
1041            # Deprecation warning for ssl_shutdown_timeout parameter
1042            warnings.warn(
1043                "The ssl_shutdown_timeout parameter is deprecated and will be removed in aiohttp 4.0",
1044                DeprecationWarning,
1045                stacklevel=2,
1046            )
1047            if (
1048                sys.version_info < (3, 11)
1049                and ssl_shutdown_timeout is not None
1050                and ssl_shutdown_timeout != 0
1051            ):
1052                warnings.warn(
1053                    f"ssl_shutdown_timeout={ssl_shutdown_timeout} is ignored on Python < 3.11; "
1054                    "only ssl_shutdown_timeout=0 is supported. The timeout will be ignored.",
1055                    RuntimeWarning,
1056                    stacklevel=2,
1057                )
1058            self._ssl_shutdown_timeout = ssl_shutdown_timeout
1059
1060    def _close(self, *, abort_ssl: bool = False) -> List[Awaitable[object]]:
1061        """Close all ongoing DNS calls."""
1062        for fut in chain.from_iterable(self._throttle_dns_futures.values()):
1063            fut.cancel()
1064
1065        waiters = super()._close(abort_ssl=abort_ssl)
1066
1067        for t in self._resolve_host_tasks:
1068            t.cancel()
1069            waiters.append(t)
1070
1071        return waiters
1072
1073    async def close(self, *, abort_ssl: bool = False) -> None:
1074        """
1075        Close all opened transports.
1076
1077        :param abort_ssl: If True, SSL connections will be aborted immediately
1078                         without performing the shutdown handshake. If False (default),
1079                         the behavior is determined by ssl_shutdown_timeout:
1080                         - If ssl_shutdown_timeout=0: connections are aborted
1081                         - If ssl_shutdown_timeout>0: graceful shutdown is performed
1082        """
1083        if self._resolver_owner:
1084            await self._resolver.close()
1085        # Use abort_ssl param if explicitly set, otherwise use ssl_shutdown_timeout default
1086        await super().close(abort_ssl=abort_ssl or self._ssl_shutdown_timeout == 0)
1087
1088    @property
1089    def family(self) -> int:
1090        """Socket family like AF_INET."""
1091        return self._family
1092
1093    @property
1094    def use_dns_cache(self) -> bool:
1095        """True if local DNS caching is enabled."""
1096        return self._use_dns_cache
1097
1098    def clear_dns_cache(
1099        self, host: Optional[str] = None, port: Optional[int] = None
1100    ) -> None:
1101        """Remove specified host/port or clear all dns local cache."""
1102        if host is not None and port is not None:
1103            self._cached_hosts.remove((host, port))
1104        elif host is not None or port is not None:
1105            raise ValueError("either both host and port or none of them are allowed")
1106        else:
1107            self._cached_hosts.clear()
1108
1109    async def _resolve_host(
1110        self, host: str, port: int, traces: Optional[Sequence["Trace"]] = None
1111    ) -> List[ResolveResult]:
1112        """Resolve host and return list of addresses."""
1113        if is_ip_address(host):
1114            return [
1115                {
1116                    "hostname": host,
1117                    "host": host,
1118                    "port": port,
1119                    "family": self._family,
1120                    "proto": 0,
1121                    "flags": 0,
1122                }
1123            ]
1124
1125        if not self._use_dns_cache:
1126
1127            if traces:
1128                for trace in traces:
1129                    await trace.send_dns_resolvehost_start(host)
1130
1131            res = await self._resolver.resolve(host, port, family=self._family)
1132
1133            if traces:
1134                for trace in traces:
1135                    await trace.send_dns_resolvehost_end(host)
1136
1137            return res
1138
1139        key = (host, port)
1140        if key in self._cached_hosts and not self._cached_hosts.expired(key):
1141            # get result early, before any await (#4014)
1142            result = self._cached_hosts.next_addrs(key)
1143
1144            if traces:
1145                for trace in traces:
1146                    await trace.send_dns_cache_hit(host)
1147            return result
1148
1149        futures: Set["asyncio.Future[None]"]
1150        #
1151        # If multiple connectors are resolving the same host, we wait
1152        # for the first one to resolve and then use the result for all of them.
1153        # We use a throttle to ensure that we only resolve the host once
1154        # and then use the result for all the waiters.
1155        #
1156        if key in self._throttle_dns_futures:
1157            # get futures early, before any await (#4014)
1158            futures = self._throttle_dns_futures[key]
1159            future: asyncio.Future[None] = self._loop.create_future()
1160            futures.add(future)
1161            if traces:
1162                for trace in traces:
1163                    await trace.send_dns_cache_hit(host)
1164            try:
1165                await future
1166            finally:
1167                futures.discard(future)
1168            return self._cached_hosts.next_addrs(key)
1169
1170        # update dict early, before any await (#4014)
1171        self._throttle_dns_futures[key] = futures = set()
1172        # In this case we need to create a task to ensure that we can shield
1173        # the task from cancellation as cancelling this lookup should not cancel
1174        # the underlying lookup or else the cancel event will get broadcast to
1175        # all the waiters across all connections.
1176        #
1177        coro = self._resolve_host_with_throttle(key, host, port, futures, traces)
1178        loop = asyncio.get_running_loop()
1179        if sys.version_info >= (3, 12):
1180            # Optimization for Python 3.12, try to send immediately
1181            resolved_host_task = asyncio.Task(coro, loop=loop, eager_start=True)
1182        else:
1183            resolved_host_task = loop.create_task(coro)
1184
1185        if not resolved_host_task.done():
1186            self._resolve_host_tasks.add(resolved_host_task)
1187            resolved_host_task.add_done_callback(self._resolve_host_tasks.discard)
1188
1189        try:
1190            return await asyncio.shield(resolved_host_task)
1191        except asyncio.CancelledError:
1192
1193            def drop_exception(fut: "asyncio.Future[List[ResolveResult]]") -> None:
1194                with suppress(Exception, asyncio.CancelledError):
1195                    fut.result()
1196
1197            resolved_host_task.add_done_callback(drop_exception)
1198            raise
1199
1200    async def _resolve_host_with_throttle(

Showing the first 1,200 of 1855 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai