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