Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tls.py422 linesDownload Raw Back to streams
1from __future__ import annotations2 3__all__ = (4    "TLSAttribute",5    "TLSConnectable",6    "TLSListener",7    "TLSStream",8)9 10import logging11import re12import ssl13import sys14from collections.abc import Callable, Mapping15from dataclasses import dataclass16from functools import wraps17from ssl import SSLContext18from typing import Any, TypeAlias, TypeVar19 20from .. import (21    BrokenResourceError,22    EndOfStream,23    aclose_forcefully,24    get_cancelled_exc_class,25    to_thread,26)27from .._core._typedattr import TypedAttributeSet, typed_attribute28from ..abc import (29    AnyByteStream,30    AnyByteStreamConnectable,31    ByteStream,32    ByteStreamConnectable,33    Listener,34    TaskGroup,35)36 37if sys.version_info >= (3, 11):38    from typing import TypeVarTuple, Unpack39else:40    from typing_extensions import TypeVarTuple, Unpack41 42if sys.version_info >= (3, 12):43    from typing import override44else:45    from typing_extensions import override46 47T_Retval = TypeVar("T_Retval")48PosArgsT = TypeVarTuple("PosArgsT")49_PCTRTT: TypeAlias = tuple[tuple[str, str], ...]50_PCTRTTT: TypeAlias = tuple[_PCTRTT, ...]51 52 53class TLSAttribute(TypedAttributeSet):54    """Contains Transport Layer Security related attributes."""55 56    #: the selected ALPN protocol57    alpn_protocol: str | None = typed_attribute()58    #: the channel binding for type ``tls-unique``59    channel_binding_tls_unique: bytes = typed_attribute()60    #: the selected cipher61    cipher: tuple[str, str, int] = typed_attribute()62    #: the peer certificate in dictionary form (see :meth:`ssl.SSLSocket.getpeercert`63    # for more information)64    peer_certificate: None | (dict[str, str | _PCTRTTT | _PCTRTT]) = typed_attribute()65    #: the peer certificate in binary form66    peer_certificate_binary: bytes | None = typed_attribute()67    #: ``True`` if this is the server side of the connection68    server_side: bool = typed_attribute()69    #: ciphers shared by the client during the TLS handshake (``None`` if this is the70    #: client side)71    shared_ciphers: list[tuple[str, str, int]] | None = typed_attribute()72    #: the :class:`~ssl.SSLObject` used for encryption73    ssl_object: ssl.SSLObject = typed_attribute()74    #: ``True`` if this stream does (and expects) a closing TLS handshake when the75    #: stream is being closed76    standard_compatible: bool = typed_attribute()77    #: the TLS protocol version (e.g. ``TLSv1.2``)78    tls_version: str = typed_attribute()79 80 81@dataclass(eq=False)82class TLSStream(ByteStream):83    """84    A stream wrapper that encrypts all sent data and decrypts received data.85 86    This class has no public initializer; use :meth:`wrap` instead.87    All extra attributes from :class:`~TLSAttribute` are supported.88 89    :var AnyByteStream transport_stream: the wrapped stream90 91    """92 93    transport_stream: AnyByteStream94    standard_compatible: bool95    _ssl_object: ssl.SSLObject96    _read_bio: ssl.MemoryBIO97    _write_bio: ssl.MemoryBIO98 99    @classmethod100    async def wrap(101        cls,102        transport_stream: AnyByteStream,103        *,104        server_side: bool | None = None,105        hostname: str | None = None,106        ssl_context: ssl.SSLContext | None = None,107        standard_compatible: bool = True,108    ) -> TLSStream:109        """110        Wrap an existing stream with Transport Layer Security.111 112        This performs a TLS handshake with the peer.113 114        :param transport_stream: a bytes-transporting stream to wrap115        :param server_side: ``True`` if this is the server side of the connection,116            ``False`` if this is the client side (if omitted, will be set to ``False``117            if ``hostname`` has been provided, ``False`` otherwise). Used only to create118            a default context when an explicit context has not been provided.119        :param hostname: host name of the peer (if host name checking is desired)120        :param ssl_context: the SSLContext object to use (if not provided, a secure121            default will be created)122        :param standard_compatible: if ``False``, skip the closing handshake when123            closing the connection, and don't raise an exception if the peer does the124            same125        :raises ~ssl.SSLError: if the TLS handshake fails126 127        """128        if server_side is None:129            server_side = not hostname130 131        if not ssl_context:132            purpose = (133                ssl.Purpose.CLIENT_AUTH if server_side else ssl.Purpose.SERVER_AUTH134            )135            ssl_context = ssl.create_default_context(purpose)136 137            # Re-enable detection of unexpected EOFs if it was disabled by Python138            if hasattr(ssl, "OP_IGNORE_UNEXPECTED_EOF"):139                ssl_context.options &= ~ssl.OP_IGNORE_UNEXPECTED_EOF140 141        bio_in = ssl.MemoryBIO()142        bio_out = ssl.MemoryBIO()143 144        # External SSLContext implementations may do blocking I/O in wrap_bio(),145        # but the standard library implementation won't146        if type(ssl_context) is ssl.SSLContext:147            ssl_object = ssl_context.wrap_bio(148                bio_in, bio_out, server_side=server_side, server_hostname=hostname149            )150        else:151            ssl_object = await to_thread.run_sync(152                ssl_context.wrap_bio,153                bio_in,154                bio_out,155                server_side,156                hostname,157                None,158            )159 160        wrapper = cls(161            transport_stream=transport_stream,162            standard_compatible=standard_compatible,163            _ssl_object=ssl_object,164            _read_bio=bio_in,165            _write_bio=bio_out,166        )167        await wrapper._call_sslobject_method(ssl_object.do_handshake)168        return wrapper169 170    async def _call_sslobject_method(171        self, func: Callable[[Unpack[PosArgsT]], T_Retval], *args: Unpack[PosArgsT]172    ) -> T_Retval:173        while True:174            try:175                result = func(*args)176            except ssl.SSLWantReadError:177                try:178                    # Flush any pending writes first179                    if self._write_bio.pending:180                        await self.transport_stream.send(self._write_bio.read())181 182                    data = await self.transport_stream.receive()183                except EndOfStream:184                    self._read_bio.write_eof()185                except OSError as exc:186                    self._read_bio.write_eof()187                    self._write_bio.write_eof()188                    raise BrokenResourceError from exc189                else:190                    self._read_bio.write(data)191            except ssl.SSLWantWriteError:192                await self.transport_stream.send(self._write_bio.read())193            except ssl.SSLSyscallError as exc:194                self._read_bio.write_eof()195                self._write_bio.write_eof()196                raise BrokenResourceError from exc197            except ssl.SSLError as exc:198                self._read_bio.write_eof()199                self._write_bio.write_eof()200                if isinstance(exc, ssl.SSLEOFError) or (201                    exc.strerror and "UNEXPECTED_EOF_WHILE_READING" in exc.strerror202                ):203                    if self.standard_compatible:204                        raise BrokenResourceError from exc205                    else:206                        raise EndOfStream from None207 208                raise209            else:210                # Flush any pending writes first211                if self._write_bio.pending:212                    await self.transport_stream.send(self._write_bio.read())213 214                return result215 216    async def unwrap(self) -> tuple[AnyByteStream, bytes]:217        """218        Does the TLS closing handshake.219 220        :return: a tuple of (wrapped byte stream, bytes left in the read buffer)221 222        """223        await self._call_sslobject_method(self._ssl_object.unwrap)224        self._read_bio.write_eof()225        self._write_bio.write_eof()226        return self.transport_stream, self._read_bio.read()227 228    async def aclose(self) -> None:229        if self.standard_compatible:230            try:231                await self.unwrap()232            except BaseException:233                await aclose_forcefully(self.transport_stream)234                raise235 236        await self.transport_stream.aclose()237 238    async def receive(self, max_bytes: int = 65536) -> bytes:239        data = await self._call_sslobject_method(self._ssl_object.read, max_bytes)240        if not data:241            raise EndOfStream242 243        return data244 245    async def send(self, item: bytes) -> None:246        await self._call_sslobject_method(self._ssl_object.write, item)247 248    async def send_eof(self) -> None:249        tls_version = self.extra(TLSAttribute.tls_version)250        match = re.match(r"TLSv(\d+)(?:\.(\d+))?", tls_version)251        if match:252            major, minor = int(match.group(1)), int(match.group(2) or 0)253            if (major, minor) < (1, 3):254                raise NotImplementedError(255                    f"send_eof() requires at least TLSv1.3; current "256                    f"session uses {tls_version}"257                )258 259        raise NotImplementedError(260            "send_eof() has not yet been implemented for TLS streams"261        )262 263    @property264    def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:265        return {266            **self.transport_stream.extra_attributes,267            TLSAttribute.alpn_protocol: self._ssl_object.selected_alpn_protocol,268            TLSAttribute.channel_binding_tls_unique: (269                self._ssl_object.get_channel_binding270            ),271            TLSAttribute.cipher: self._ssl_object.cipher,272            TLSAttribute.peer_certificate: lambda: self._ssl_object.getpeercert(False),273            TLSAttribute.peer_certificate_binary: lambda: self._ssl_object.getpeercert(274                True275            ),276            TLSAttribute.server_side: lambda: self._ssl_object.server_side,277            TLSAttribute.shared_ciphers: lambda: (278                self._ssl_object.shared_ciphers()279                if self._ssl_object.server_side280                else None281            ),282            TLSAttribute.standard_compatible: lambda: self.standard_compatible,283            TLSAttribute.ssl_object: lambda: self._ssl_object,284            TLSAttribute.tls_version: self._ssl_object.version,285        }286 287 288@dataclass(eq=False)289class TLSListener(Listener[TLSStream]):290    """291    A convenience listener that wraps another listener and auto-negotiates a TLS session292    on every accepted connection.293 294    If the TLS handshake times out or raises an exception,295    :meth:`handle_handshake_error` is called to do whatever post-mortem processing is296    deemed necessary.297 298    Supports only the :attr:`~TLSAttribute.standard_compatible` extra attribute.299 300    :param Listener listener: the listener to wrap301    :param ssl_context: the SSL context object302    :param standard_compatible: a flag passed through to :meth:`TLSStream.wrap`303    :param handshake_timeout: time limit for the TLS handshake304        (passed to :func:`~anyio.fail_after`)305    """306 307    listener: Listener[Any]308    ssl_context: ssl.SSLContext309    standard_compatible: bool = True310    handshake_timeout: float = 30311 312    @staticmethod313    async def handle_handshake_error(exc: BaseException, stream: AnyByteStream) -> None:314        """315        Handle an exception raised during the TLS handshake.316 317        This method does 3 things:318 319        #. Forcefully closes the original stream320        #. Logs the exception (unless it was a cancellation exception) using the321           ``anyio.streams.tls`` logger322        #. Reraises the exception if it was a base exception or a cancellation exception323 324        :param exc: the exception325        :param stream: the original stream326 327        """328        await aclose_forcefully(stream)329 330        # Log all except cancellation exceptions331        if not isinstance(exc, get_cancelled_exc_class()):332            # CPython (as of 3.11.5) returns incorrect `sys.exc_info()` here when using333            # any asyncio implementation, so we explicitly pass the exception to log334            # (https://github.com/python/cpython/issues/108668). Trio does not have this335            # issue because it works around the CPython bug.336            logging.getLogger(__name__).exception(337                "Error during TLS handshake", exc_info=exc338            )339 340        # Only reraise base exceptions and cancellation exceptions341        if not isinstance(exc, Exception) or isinstance(exc, get_cancelled_exc_class()):342            raise343 344    async def serve(345        self,346        handler: Callable[[TLSStream], Any],347        task_group: TaskGroup | None = None,348    ) -> None:349        @wraps(handler)350        async def handler_wrapper(stream: AnyByteStream) -> None:351            from .. import fail_after352 353            try:354                with fail_after(self.handshake_timeout):355                    wrapped_stream = await TLSStream.wrap(356                        stream,357                        ssl_context=self.ssl_context,358                        standard_compatible=self.standard_compatible,359                    )360            except BaseException as exc:361                await self.handle_handshake_error(exc, stream)362            else:363                await handler(wrapped_stream)364 365        await self.listener.serve(handler_wrapper, task_group)366 367    async def aclose(self) -> None:368        await self.listener.aclose()369 370    @property371    def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:372        return {373            TLSAttribute.standard_compatible: lambda: self.standard_compatible,374        }375 376 377class TLSConnectable(ByteStreamConnectable):378    """379    Wraps another connectable and does TLS negotiation after a successful connection.380 381    :param connectable: the connectable to wrap382    :param hostname: host name of the server (if host name checking is desired)383    :param ssl_context: the SSLContext object to use (if not provided, a secure default384        will be created)385    :param standard_compatible: if ``False``, skip the closing handshake when closing386        the connection, and don't raise an exception if the server does the same387    """388 389    def __init__(390        self,391        connectable: AnyByteStreamConnectable,392        *,393        hostname: str | None = None,394        ssl_context: ssl.SSLContext | None = None,395        standard_compatible: bool = True,396    ) -> None:397        self.connectable = connectable398        self.ssl_context: SSLContext = ssl_context or ssl.create_default_context(399            ssl.Purpose.SERVER_AUTH400        )401        if not isinstance(self.ssl_context, ssl.SSLContext):402            raise TypeError(403                "ssl_context must be an instance of ssl.SSLContext, not "404                f"{type(self.ssl_context).__name__}"405            )406        self.hostname = hostname407        self.standard_compatible = standard_compatible408 409    @override410    async def connect(self) -> TLSStream:411        stream = await self.connectable.connect()412        try:413            return await TLSStream.wrap(414                stream,415                hostname=self.hostname,416                ssl_context=self.ssl_context,417                standard_compatible=self.standard_compatible,418            )419        except BaseException:420            await aclose_forcefully(stream)421            raise422 
codekingpro/portable-devtools · Team Ai