codekingpro/portable-devtools
114k
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(
