From 5659901428dd1241688f1090dfa2c374c51ad76c Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Thu, 8 Oct 2026 19:26:23 +0000 Subject: [PATCH 1/4] PYTHON-6154 Extract KMS connect helpers into dedicated modules Move the KMS connect plumbing out of the encryption modules into dedicated private modules, prep for the CSFLE HTTP proxy KMS connect helpers (PYTHON-6147): - pymongo/_kms_connect.py holds KMSConnectContext, the callback type aliases, and _close_rejected_kms_socket, shared by both APIs. - pymongo/{asynchronous,synchronous}/_kms_connect.py hold _connect_kms and _KMS_CONNECT_TIMEOUT; the synchronous module is generated by synchro. - encryption_options.py re-exports the public names unchanged. - test/asynchronous/test_kms_connect.py is expanded to cover the new helpers and test/test_kms_connect.py is its synchro-generated mirror. --- pymongo/_kms_connect.py | 79 ++++ pymongo/asynchronous/_kms_connect.py | 146 +++++++ pymongo/asynchronous/encryption.py | 119 +----- pymongo/encryption_options.py | 41 +- pymongo/synchronous/_kms_connect.py | 146 +++++++ pymongo/synchronous/encryption.py | 120 +----- test/asynchronous/test_kms_connect.py | 535 +++++++++++++++++++++++++- test/test_kms_connect.py | 533 ++++++++++++++++++++++++- 8 files changed, 1443 insertions(+), 276 deletions(-) create mode 100644 pymongo/_kms_connect.py create mode 100644 pymongo/asynchronous/_kms_connect.py create mode 100644 pymongo/synchronous/_kms_connect.py diff --git a/pymongo/_kms_connect.py b/pymongo/_kms_connect.py new file mode 100644 index 0000000000..50cd3d73f0 --- /dev/null +++ b/pymongo/_kms_connect.py @@ -0,0 +1,79 @@ +# Copyright 2026-present MongoDB, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""KMS connection helpers, shared by both the synchronous and asynchronous APIs. + +Holds the ``KMSConnectContext`` passed to ``kms_connect_callback`` and the +callback type aliases. The per-API ``_connect_kms`` that connects to the KMS +host and performs the TLS handshake lives in +``pymongo.asynchronous._kms_connect`` and its generated synchronous mirror. +""" + +from __future__ import annotations + +import contextlib +import socket +from collections.abc import Awaitable +from dataclasses import dataclass +from typing import Any, Callable + + +@dataclass(frozen=True) +class KMSConnectContext: + """Information about a pending KMS connection. + + Passed to ``kms_connect_callback``, which must return a plain, unwrapped + :class:`socket.socket`. The driver performs the KMS TLS handshake over it, + verifying against ``host`` rather than the peer actually reached. + + :param host: Hostname of the KMS server, and the TLS verification target. + :param port: Port of the KMS server. + :param timeout: Seconds allowed for the connection: the default KMS + connect timeout, capped by the remaining time of an active operation + timeout (``timeoutMS``). + + .. note:: ``timeoutMS`` does not constrain KMS requests for explicit + encryption, so ``timeout`` is always the default there. Automatic + encryption passes the remaining budget. This deviates from the Client + Side Operations Timeout specification; see PYTHON-6037. + + .. versionadded:: 4.19 + """ + + host: str + port: int + timeout: float + + +# A callback that opens a connection to a KMS host. +AsyncKMSConnectCallback = Callable[[KMSConnectContext], Awaitable[socket.socket]] +KMSConnectCallback = Callable[[KMSConnectContext], socket.socket] + + +def _close_rejected_kms_socket(obj: Any) -> None: + """Close a callback return value that failed validation, best effort. + + ``_connect_kms`` raises on a contract violation instead of returning the + value, so no caller ever takes ownership of it. Close it here, tolerating + non-socket values and ``close()`` failures. + """ + close = getattr(obj, "close", None) + if callable(close): + with contextlib.suppress(Exception): + close() + + +# Sphinx documents this class under pymongo.encryption_options, the public +# import path, so the definition must claim that module name. +KMSConnectContext.__module__ = "pymongo.encryption_options" diff --git a/pymongo/asynchronous/_kms_connect.py b/pymongo/asynchronous/_kms_connect.py new file mode 100644 index 0000000000..176710e67d --- /dev/null +++ b/pymongo/asynchronous/_kms_connect.py @@ -0,0 +1,146 @@ +# Copyright 2026-present MongoDB, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""KMS connection for the asynchronous API. + +Connects to a KMS host, directly or through a ``kms_connect_callback``, and +performs the KMS TLS handshake over the connection. The synchronous mirror of +this module is generated by synchro. The helpers shared by both APIs live in +``pymongo._kms_connect``. +""" + +from __future__ import annotations + +import asyncio +import inspect +import socket +import ssl +from typing import TYPE_CHECKING, Any, Optional, Union, cast + +from pymongo import _csot +from pymongo._kms_connect import ( + AsyncKMSConnectCallback, + KMSConnectContext, + _close_rejected_kms_socket, +) +from pymongo.common import CONNECT_TIMEOUT +from pymongo.errors import ConfigurationError +from pymongo.helpers_shared import _get_timeout_details +from pymongo.pool_options import PoolOptions +from pymongo.pool_shared import ( + _async_configured_socket, + _async_wrap_socket_tls, + _close_late_socket, + _raise_connection_failure, +) + +if TYPE_CHECKING: + from pymongo.pyopenssl_context import _sslConn + from pymongo.typings import _Address + + +_IS_SYNC = False + +_KMS_CONNECT_TIMEOUT = CONNECT_TIMEOUT # CDRIVER-3262 redefined this value to CONNECT_TIMEOUT + + +async def _connect_kms( + address: _Address, + opts: PoolOptions, + kms_connect_callback: Optional[AsyncKMSConnectCallback], + timeout: float, +) -> Union[socket.socket, _sslConn]: + """Connect to a KMS host and perform the TLS handshake over the socket. + + Uses ``kms_connect_callback`` when one is provided, otherwise connects + directly, and always verifies against ``address`` (the KMS host). + """ + if kms_connect_callback is None: + try: + return await _async_configured_socket(address, opts) + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + + # TLS targets address, not the peer, so verification follows the KMS host. + # A plain callable would block the event loop before we could reject it, + # so check the callback first. + if not _IS_SYNC: + callback_any: Any = kms_connect_callback + is_coro = inspect.iscoroutinefunction(callback_any) + if not is_coro and callable(callback_any): + is_coro = inspect.iscoroutinefunction(callback_any.__call__) + if not is_coro: + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async API." + ) + # Any: the sync version's callback returns a plain socket. + result: Any = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + remaining = _csot.remaining() + if remaining is None or _IS_SYNC: + # The synchronous API cannot interrupt a callback that has started + # running; honoring the deadline is the callback's contract there. + sock = await result + else: + # CSOT is cooperative: a callback that ignores the timeout could block + # past the deadline. Shield the task so stopping the wait does not + # cancel it mid-flight, and close any socket it yields later. + task = asyncio.ensure_future(result) + try: + sock = await asyncio.wait_for(asyncio.shield(task), remaining) + except asyncio.CancelledError: + task.add_done_callback(_close_late_socket) + raise + except asyncio.TimeoutError: + task.add_done_callback(_close_late_socket) + _raise_connection_failure( + address, + socket.timeout("timed out"), + timeout_details=_get_timeout_details(opts), + ) + if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a connected, unwrapped " + f"socket.socket, not {type(sock)}." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. + try: + sock.getpeername() + except OSError: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return an already connected socket." + ) from None + if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a stream socket, not a datagram one." + ) + # The callback may have consumed much of the CSOT budget, and wrapping + # resets the socket timeout, so recompute the remaining time here and for + # the KMS request that follows. + sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + try: + conn = await _async_wrap_socket_tls(sock, address, opts) + except asyncio.CancelledError: + # The executor may still be wrapping the socket; close it so a TLS + # proxy's relay threads wind down instead of leaking. + sock.close() + raise + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + return conn diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 4ae63716e3..955529fa97 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -17,11 +17,8 @@ from __future__ import annotations import asyncio -import contextlib import functools -import inspect import socket -import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -56,16 +53,15 @@ from bson.codec_options import CodecOptions from bson.raw_bson import DEFAULT_RAW_BSON_OPTIONS, RawBSONDocument, _inflate_bson from pymongo import _csot, _op_id +from pymongo._kms_connect import AsyncKMSConnectCallback +from pymongo.asynchronous._kms_connect import _KMS_CONNECT_TIMEOUT, _connect_kms from pymongo.asynchronous.collection import AsyncCollection from pymongo.asynchronous.cursor import AsyncCursor from pymongo.asynchronous.database import AsyncDatabase from pymongo.asynchronous.mongo_client import AsyncMongoClient -from pymongo.common import CONNECT_TIMEOUT from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( - AsyncKMSConnectCallback, AutoEncryptionOpts, - KMSConnectContext, RangeOpts, StringOpts, # Re-exported for backwards compatibility: TextOpts is deprecated but must @@ -94,9 +90,6 @@ from pymongo.operations import UpdateOne from pymongo.pool_options import PoolOptions from pymongo.pool_shared import ( - _async_configured_socket, - _async_wrap_socket_tls, - _close_late_socket, _raise_connection_failure, ) from pymongo.read_concern import ReadConcern @@ -109,14 +102,10 @@ if TYPE_CHECKING: from pymongocrypt.mongocrypt import MongoCryptKmsContext - from pymongo.pyopenssl_context import _sslConn - from pymongo.typings import _Address - _IS_SYNC = False _HTTPS_PORT = 443 -_KMS_CONNECT_TIMEOUT = CONNECT_TIMEOUT # CDRIVER-3262 redefined this value to CONNECT_TIMEOUT _MONGOCRYPTD_TIMEOUT_MS = 10000 _DATA_KEY_OPTS: CodecOptions[dict[str, Any]] = CodecOptions( @@ -127,110 +116,6 @@ _KEY_VAULT_OPTS = CodecOptions(document_class=RawBSONDocument) -def _close_rejected_kms_socket(obj: Any) -> None: - """Close a callback return value that failed validation, best effort. - - ``_connect_kms`` raises on a contract violation instead of returning the - value, so no caller ever takes ownership of it. Close it here, tolerating - non-socket values and ``close()`` failures. - """ - close = getattr(obj, "close", None) - if callable(close): - with contextlib.suppress(Exception): - close() - - -async def _connect_kms( - address: _Address, - opts: PoolOptions, - kms_connect_callback: Optional[AsyncKMSConnectCallback], - timeout: float, -) -> Union[socket.socket, _sslConn]: - """Connect to a KMS host and perform the TLS handshake over the socket. - - Uses ``kms_connect_callback`` when one is provided, otherwise connects - directly, and always verifies against ``address`` (the KMS host). - """ - if kms_connect_callback is None: - try: - return await _async_configured_socket(address, opts) - except Exception as exc: - _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - - # TLS targets address, not the peer, so verification follows the KMS host. - # A plain callable would block the event loop before we could reject it, - # so check the callback first. - if not _IS_SYNC: - callback_any: Any = kms_connect_callback - is_coro = inspect.iscoroutinefunction(callback_any) - if not is_coro and callable(callback_any): - is_coro = inspect.iscoroutinefunction(callback_any.__call__) - if not is_coro: - raise ConfigurationError( - "kms_connect_callback must be a coroutine function for the async API." - ) - # Any: the sync version's callback returns a plain socket. - result: Any = kms_connect_callback( - KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) - ) - remaining = _csot.remaining() - if remaining is None or _IS_SYNC: - # The synchronous API cannot interrupt a callback that has started - # running; honoring the deadline is the callback's contract there. - sock = await result - else: - # CSOT is cooperative: a callback that ignores the timeout could block - # past the deadline. Shield the task so stopping the wait does not - # cancel it mid-flight, and close any socket it yields later. - task = asyncio.ensure_future(result) - try: - sock = await asyncio.wait_for(asyncio.shield(task), remaining) - except asyncio.CancelledError: - task.add_done_callback(_close_late_socket) - raise - except asyncio.TimeoutError: - task.add_done_callback(_close_late_socket) - _raise_connection_failure( - address, - socket.timeout("timed out"), - timeout_details=_get_timeout_details(opts), - ) - if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return a connected, unwrapped " - f"socket.socket, not {type(sock)}." - ) - # wrap_socket refuses a non-blocking socket, so normalize the mode here. - try: - sock.getpeername() - except OSError: - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return an already connected socket." - ) from None - if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return a stream socket, not a datagram one." - ) - # The callback may have consumed much of the CSOT budget, and wrapping - # resets the socket timeout, so recompute the remaining time here and for - # the KMS request that follows. - sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) - try: - conn = await _async_wrap_socket_tls(sock, address, opts) - except asyncio.CancelledError: - # The executor may still be wrapping the socket; close it so a TLS - # proxy's relay threads wind down instead of leaking. - sock.close() - raise - except Exception as exc: - _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) - return conn - - class _EncryptionIO(AsyncMongoCryptCallback): # type: ignore[misc] def __init__( self, diff --git a/pymongo/encryption_options.py b/pymongo/encryption_options.py index b1d13700da..41e200eefc 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -19,10 +19,8 @@ from __future__ import annotations -import socket import warnings -from collections.abc import Awaitable, Mapping -from dataclasses import dataclass +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Callable, Optional, TypedDict from pymongo.uri_parser_shared import _parse_kms_tls_options @@ -37,6 +35,11 @@ except ImportError: _HAVE_PYMONGOCRYPT = False from bson import int64 +from pymongo._kms_connect import ( # noqa: F401 + AsyncKMSConnectCallback, + KMSConnectCallback, + KMSConnectContext, +) from pymongo.common import check_for_min_version, validate_is_mapping from pymongo.errors import ConfigurationError @@ -57,38 +60,6 @@ def check_min_pymongocrypt() -> None: ) -@dataclass(frozen=True) -class KMSConnectContext: - """Information about a pending KMS connection. - - Passed to ``kms_connect_callback``, which must return a plain, unwrapped - :class:`socket.socket`. The driver performs the KMS TLS handshake over it, - verifying against ``host`` rather than the peer actually reached. - - :param host: Hostname of the KMS server, and the TLS verification target. - :param port: Port of the KMS server. - :param timeout: Seconds allowed for the connection: the default KMS - connect timeout, capped by the remaining time of an active operation - timeout (``timeoutMS``). - - .. note:: ``timeoutMS`` does not constrain KMS requests for explicit - encryption, so ``timeout`` is always the default there. Automatic - encryption passes the remaining budget. This deviates from the Client - Side Operations Timeout specification; see PYTHON-6037. - - .. versionadded:: 4.19 - """ - - host: str - port: int - timeout: float - - -# A callback that opens a connection to a KMS host. -AsyncKMSConnectCallback = Callable[[KMSConnectContext], Awaitable[socket.socket]] -KMSConnectCallback = Callable[[KMSConnectContext], socket.socket] - - class AutoEncryptionOpts: """Options to configure automatic client-side field level encryption.""" diff --git a/pymongo/synchronous/_kms_connect.py b/pymongo/synchronous/_kms_connect.py new file mode 100644 index 0000000000..bd453e9114 --- /dev/null +++ b/pymongo/synchronous/_kms_connect.py @@ -0,0 +1,146 @@ +# Copyright 2026-present MongoDB, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""KMS connection for the synchronous API. + +Connects to a KMS host, directly or through a ``kms_connect_callback``, and +performs the KMS TLS handshake over the connection. The synchronous mirror of +this module is generated by synchro. The helpers shared by both APIs live in +``pymongo._kms_connect``. +""" + +from __future__ import annotations + +import asyncio +import inspect +import socket +import ssl +from typing import TYPE_CHECKING, Any, Optional, Union, cast + +from pymongo import _csot +from pymongo._kms_connect import ( + KMSConnectCallback, + KMSConnectContext, + _close_rejected_kms_socket, +) +from pymongo.common import CONNECT_TIMEOUT +from pymongo.errors import ConfigurationError +from pymongo.helpers_shared import _get_timeout_details +from pymongo.pool_options import PoolOptions +from pymongo.pool_shared import ( + _close_late_socket, + _configured_socket, + _raise_connection_failure, + _wrap_socket_tls, +) + +if TYPE_CHECKING: + from pymongo.pyopenssl_context import _sslConn + from pymongo.typings import _Address + + +_IS_SYNC = True + +_KMS_CONNECT_TIMEOUT = CONNECT_TIMEOUT # CDRIVER-3262 redefined this value to CONNECT_TIMEOUT + + +def _connect_kms( + address: _Address, + opts: PoolOptions, + kms_connect_callback: Optional[KMSConnectCallback], + timeout: float, +) -> Union[socket.socket, _sslConn]: + """Connect to a KMS host and perform the TLS handshake over the socket. + + Uses ``kms_connect_callback`` when one is provided, otherwise connects + directly, and always verifies against ``address`` (the KMS host). + """ + if kms_connect_callback is None: + try: + return _configured_socket(address, opts) + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + + # TLS targets address, not the peer, so verification follows the KMS host. + # A plain callable would block the event loop before we could reject it, + # so check the callback first. + if not _IS_SYNC: + callback_any: Any = kms_connect_callback + is_coro = inspect.iscoroutinefunction(callback_any) + if not is_coro and callable(callback_any): + is_coro = inspect.iscoroutinefunction(callback_any.__call__) + if not is_coro: + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async API." + ) + # Any: the sync version's callback returns a plain socket. + result: Any = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + remaining = _csot.remaining() + if remaining is None or _IS_SYNC: + # The synchronous API cannot interrupt a callback that has started + # running; honoring the deadline is the callback's contract there. + sock = result + else: + # CSOT is cooperative: a callback that ignores the timeout could block + # past the deadline. Shield the task so stopping the wait does not + # cancel it mid-flight, and close any socket it yields later. + task = asyncio.ensure_future(result) + try: + sock = asyncio.wait_for(asyncio.shield(task), remaining) + except asyncio.CancelledError: + task.add_done_callback(_close_late_socket) + raise + except asyncio.TimeoutError: + task.add_done_callback(_close_late_socket) + _raise_connection_failure( + address, + socket.timeout("timed out"), + timeout_details=_get_timeout_details(opts), + ) + if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a connected, unwrapped " + f"socket.socket, not {type(sock)}." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. + try: + sock.getpeername() + except OSError: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return an already connected socket." + ) from None + if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a stream socket, not a datagram one." + ) + # The callback may have consumed much of the CSOT budget, and wrapping + # resets the socket timeout, so recompute the remaining time here and for + # the KMS request that follows. + sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + try: + conn = _wrap_socket_tls(sock, address, opts) + except asyncio.CancelledError: + # The executor may still be wrapping the socket; close it so a TLS + # proxy's relay threads wind down instead of leaking. + sock.close() + raise + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + return conn diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index ac2f710382..9d0b12c683 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -16,12 +16,8 @@ from __future__ import annotations -import asyncio -import contextlib import functools -import inspect import socket -import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -56,12 +52,10 @@ from bson.codec_options import CodecOptions from bson.raw_bson import DEFAULT_RAW_BSON_OPTIONS, RawBSONDocument, _inflate_bson from pymongo import _csot, _op_id -from pymongo.common import CONNECT_TIMEOUT +from pymongo._kms_connect import KMSConnectCallback from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( AutoEncryptionOpts, - KMSConnectCallback, - KMSConnectContext, RangeOpts, StringOpts, # Re-exported for backwards compatibility: TextOpts is deprecated but must @@ -90,14 +84,12 @@ from pymongo.operations import UpdateOne from pymongo.pool_options import PoolOptions from pymongo.pool_shared import ( - _close_late_socket, - _configured_socket, _raise_connection_failure, - _wrap_socket_tls, ) from pymongo.read_concern import ReadConcern from pymongo.results import DeleteResult from pymongo.ssl_support import BLOCKING_IO_ERRORS, get_ssl_context +from pymongo.synchronous._kms_connect import _KMS_CONNECT_TIMEOUT, _connect_kms from pymongo.synchronous.collection import Collection from pymongo.synchronous.cursor import Cursor from pymongo.synchronous.database import Database @@ -109,14 +101,10 @@ if TYPE_CHECKING: from pymongocrypt.mongocrypt import MongoCryptKmsContext - from pymongo.pyopenssl_context import _sslConn - from pymongo.typings import _Address - _IS_SYNC = True _HTTPS_PORT = 443 -_KMS_CONNECT_TIMEOUT = CONNECT_TIMEOUT # CDRIVER-3262 redefined this value to CONNECT_TIMEOUT _MONGOCRYPTD_TIMEOUT_MS = 10000 _DATA_KEY_OPTS: CodecOptions[dict[str, Any]] = CodecOptions( @@ -127,110 +115,6 @@ _KEY_VAULT_OPTS = CodecOptions(document_class=RawBSONDocument) -def _close_rejected_kms_socket(obj: Any) -> None: - """Close a callback return value that failed validation, best effort. - - ``_connect_kms`` raises on a contract violation instead of returning the - value, so no caller ever takes ownership of it. Close it here, tolerating - non-socket values and ``close()`` failures. - """ - close = getattr(obj, "close", None) - if callable(close): - with contextlib.suppress(Exception): - close() - - -def _connect_kms( - address: _Address, - opts: PoolOptions, - kms_connect_callback: Optional[KMSConnectCallback], - timeout: float, -) -> Union[socket.socket, _sslConn]: - """Connect to a KMS host and perform the TLS handshake over the socket. - - Uses ``kms_connect_callback`` when one is provided, otherwise connects - directly, and always verifies against ``address`` (the KMS host). - """ - if kms_connect_callback is None: - try: - return _configured_socket(address, opts) - except Exception as exc: - _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - - # TLS targets address, not the peer, so verification follows the KMS host. - # A plain callable would block the event loop before we could reject it, - # so check the callback first. - if not _IS_SYNC: - callback_any: Any = kms_connect_callback - is_coro = inspect.iscoroutinefunction(callback_any) - if not is_coro and callable(callback_any): - is_coro = inspect.iscoroutinefunction(callback_any.__call__) - if not is_coro: - raise ConfigurationError( - "kms_connect_callback must be a coroutine function for the async API." - ) - # Any: the sync version's callback returns a plain socket. - result: Any = kms_connect_callback( - KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) - ) - remaining = _csot.remaining() - if remaining is None or _IS_SYNC: - # The synchronous API cannot interrupt a callback that has started - # running; honoring the deadline is the callback's contract there. - sock = result - else: - # CSOT is cooperative: a callback that ignores the timeout could block - # past the deadline. Shield the task so stopping the wait does not - # cancel it mid-flight, and close any socket it yields later. - task = asyncio.ensure_future(result) - try: - sock = asyncio.wait_for(asyncio.shield(task), remaining) - except asyncio.CancelledError: - task.add_done_callback(_close_late_socket) - raise - except asyncio.TimeoutError: - task.add_done_callback(_close_late_socket) - _raise_connection_failure( - address, - socket.timeout("timed out"), - timeout_details=_get_timeout_details(opts), - ) - if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return a connected, unwrapped " - f"socket.socket, not {type(sock)}." - ) - # wrap_socket refuses a non-blocking socket, so normalize the mode here. - try: - sock.getpeername() - except OSError: - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return an already connected socket." - ) from None - if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: - _close_rejected_kms_socket(sock) - raise ConfigurationError( - "kms_connect_callback must return a stream socket, not a datagram one." - ) - # The callback may have consumed much of the CSOT budget, and wrapping - # resets the socket timeout, so recompute the remaining time here and for - # the KMS request that follows. - sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) - try: - conn = _wrap_socket_tls(sock, address, opts) - except asyncio.CancelledError: - # The executor may still be wrapping the socket; close it so a TLS - # proxy's relay threads wind down instead of leaking. - sock.close() - raise - except Exception as exc: - _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) - return conn - - class _EncryptionIO(MongoCryptCallback): # type: ignore[misc] def __init__( self, diff --git a/test/asynchronous/test_kms_connect.py b/test/asynchronous/test_kms_connect.py index 6a577b90ee..a1a6ee2205 100644 --- a/test/asynchronous/test_kms_connect.py +++ b/test/asynchronous/test_kms_connect.py @@ -1,9 +1,10 @@ -"""Tests for the KMS connect callback.""" +"""Tests for the KMS connect callback and HTTP proxy support.""" from __future__ import annotations import asyncio import dataclasses +import http.client import os import socket import ssl @@ -11,11 +12,13 @@ import time import unittest from asyncio.trsock import TransportSocket +from typing import Any from unittest import mock import pytest import pymongo +from bson.binary import Binary from pymongo.asynchronous.encryption import ( AsyncClientEncryption, _connect_kms, @@ -23,15 +26,17 @@ _wrap_encryption_errors, ) from pymongo.encryption_options import ( + AsyncHTTPProxyKMSConnect, AutoEncryptionOpts, + HTTPProxyKMSConnect, KMSConnectContext, ) from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout from pymongo.pool_options import PoolOptions from pymongo.ssl_support import get_ssl_context from test.asynchronous import AsyncPyMongoTestCase -from test.asynchronous.test_encryption import OPTS -from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM +from test.asynchronous.test_encryption import OPTS, AsyncEncryptionIntegrationTest +from test.helpers_shared import AWS_CREDS, CA_PEM, CERT_PATH, CLIENT_PEM _IS_SYNC = False @@ -39,6 +44,15 @@ _KMS_ADDRESS = ("kms.example.com", 443) +KMS_PROXY_HOST = "127.0.0.1" +KMS_PROXY_PORT = 9004 +KMS_TLS_PROXY_PORT = 9005 + +AWS_MASTER_KEY = { + "region": "us-east-1", + "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", +} + def _tls_server_context(cert=CLIENT_PEM): ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) @@ -69,6 +83,31 @@ async def callback(context): return callback +def _kms_context(host="kms.example.com", port=443, timeout=10): + """A KMSConnectContext with the defaults used throughout these tests.""" + return KMSConnectContext(host=host, port=port, timeout=timeout) + + +def _read_http_request(conn): + """Read until the blank line that ends a CONNECT request, or None on EOF.""" + request = b"" + while b"\r\n\r\n" not in request: + chunk = conn.recv(4096) + if not chunk: + return None + request += chunk + return request + + +def _insecure_client_context(): + # PYTHON-5040 tracks re-enabling verification: the evergreen-tools CA + # lacks an Authority Key Identifier newer OpenSSL requires. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + return ctx + + class TestKmsConnectCallbackUnit(AsyncPyMongoTestCase): """Contract checks for kms_connect_callback that need no KMS server.""" @@ -89,6 +128,63 @@ def _socketpair(self): self.addCleanup(right.close) return left, right + def _start_proxy(self, handler, backlog=1): + """Serve each accepted connection with ``handler(conn)`` in a daemon thread.""" + listener = self._listen(backlog) + + def serve(): + for _ in range(backlog): + try: + conn, _ = listener.accept() + except OSError: + return + try: + handler(conn) + except OSError: + pass + finally: + conn.close() + + threading.Thread(target=serve, daemon=True).start() + return listener.getsockname() + + def _record_and_reply(self, accepted, reply): + """A proxy that records each CONNECT request, replies ``reply``, and closes.""" + + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + accepted.append(request) + conn.sendall(reply) + + return self._start_proxy(handler) + + def _tls_echo_proxy(self, delay=0): + """A TLS CONNECT proxy that replies 200, then echoes one tunneled read.""" + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + request = _read_http_request(tls) + if request is None: + return + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # The tunneled peer speaks only after the client does, as a TLS + # server would; ``delay`` lets the reply outlast the CONNECT deadline. + if delay: + time.sleep(delay) + tls.sendall(b"echo:" + tls.recv(64)) + tls.close() + + return self._start_proxy(handler) + + async def _echo_over_tunnel(self, sock): + sock.settimeout(10) + await _run_blocking(sock.sendall, b"ping") + data = await _run_blocking(sock.recv, 64) + self.assertEqual(data, b"echo:ping") + async def test_init_kms_connect_callback(self): opts = AutoEncryptionOpts({}, "k.d") self.assertIsNone(opts._kms_connect_callback) @@ -242,6 +338,259 @@ async def test_cancelled_tls_wrap_closes_late_socket(self): _close_late_socket(future) self.assertEqual(left.fileno(), -1) + async def test_http_proxy_helper_tunnels_and_reports_refusal(self): + # Covers the CONNECT handshake without KMS credentials. + accepted: list[bytes] = [] + context = _kms_context() + + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + sock = await AsyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n" + ) + with self.assertRaisesRegex(OSError, "refused CONNECT"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + + # Any 2xx status is a successful tunnel, not just HTTP/1.1 200. + host, port = self._record_and_reply( + accepted, b"HTTP/1.0 200 Connection Established\r\n\r\n" + ) + sock = await AsyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + + # A status code must be exactly three digits, with no zero padding. + for reply in (b"HTTP/1.1 2000 Evil\r\n\r\n", b"HTTP/1.1 00200 Evil\r\n\r\n"): + host, port = self._record_and_reply(accepted, reply) + with self.assertRaisesRegex(OSError, "refused CONNECT"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + + async def test_control_characters_in_kms_host_are_rejected(self): + # Reject CR/LF in the configurable host before it reaches CONNECT. + callback = AsyncHTTPProxyKMSConnect("proxy.example.com", 8080) + context = _kms_context(host="kms.example.com\r\nX-Injected: 1") + with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): + await callback(context) + # Whitespace would split the request line into extra tokens. + context = _kms_context(host="kms.example.com ") + with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): + await callback(context) + + async def test_http_proxy_helper_sends_custom_headers(self): + # Extra CONNECT headers reach the proxy verbatim. + accepted: list[bytes] = [] + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Trace-Id": "abc123"} + sock = await AsyncHTTPProxyKMSConnect(host, port, headers=headers)(_kms_context()) + self.addCleanup(sock.close) + request = accepted[0] + self.assertEqual(request.split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + self.assertIn(b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n", request) + self.assertIn(b"\r\nX-Trace-Id: abc123\r\n", request) + self.assertEqual(request.count(b"\r\nHost: "), 1) + + async def test_http_proxy_helper_authenticates_to_the_proxy(self): + # The motivating case: 407 without credentials, 200 with them. + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + if b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request: + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + else: + conn.sendall(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") + + host, port = self._start_proxy(handler, backlog=2) + context = _kms_context() + with self.assertRaisesRegex(OSError, "refused CONNECT"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz"} + sock = await AsyncHTTPProxyKMSConnect(host, port, headers=headers)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + + async def test_http_proxy_helper_rejects_bad_headers(self): + for headers in [ + {"Bad\r\nName": "x"}, + {"Bad Name": "x"}, + {"Bad\tName": "x"}, + {"X-Ok": "ok\r\nInjected: 1"}, + {"Host": "evil.example.com"}, + {"host": "evil.example.com"}, + {"": "x"}, + {"Bad:Name": "x"}, + ]: + with self.assertRaisesRegex(ConfigurationError, "proxy header|Host CONNECT header"): + AsyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) + + for headers in [{1: "x"}, {"X-Ok": 1}, {None: "x"}, {"X-Ok": None}]: + with self.assertRaisesRegex(TypeError, "must be strings"): + AsyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) + + async def test_http_proxy_helper_accepts_legal_header_values(self): + # Colons and spaces are legal in values (e.g. auth schemes); only + # CR/LF would let a value inject a request line. + callback = AsyncHTTPProxyKMSConnect( + "proxy.example.com", + 8080, + headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, + ) + self.assertEqual( + callback.headers, + {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, + ) + + async def test_tls_proxy_helper_bridges_the_tunnel(self): + # Covers the TLS-proxy path and the socketpair relay without KMS creds. + host, port = self._tls_echo_proxy() + sock = await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())( + _kms_context() + ) + self.addCleanup(sock.close) + await self._echo_over_tunnel(sock) + + async def test_bridge_does_not_inherit_the_connect_deadline(self): + # The relay must outlast the much shorter CONNECT deadline. + host, port = self._tls_echo_proxy(delay=3.0) + sock = await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())( + _kms_context(timeout=2.0) + ) + self.addCleanup(sock.close) + await self._echo_over_tunnel(sock) + + async def test_proxy_closing_before_connect_reply_raises(self): + def handler(conn): + # Read the CONNECT request, then hang up without replying. + conn.recv(4096) + + host, port = self._start_proxy(handler) + with self.assertRaisesRegex(OSError, "proxy closed the connection"): + await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) + + async def test_cancelled_proxy_connect_closes_the_late_socket(self): + # A cancelled connect must close the socket the executor thread + # produces after the cancellation. + if _IS_SYNC: + raise unittest.SkipTest("cancellation is an async-only behavior") + + requested = threading.Event() + reply = threading.Event() + + def handler(conn): + conn.recv(4096) + requested.set() + if not reply.wait(10): + return + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # Keep the connection open so the tunnel can complete its reads. + time.sleep(0.1) + + host, port = self._start_proxy(handler) + + tunneled: list[socket.socket] = [] + original_tunnel = HTTPProxyKMSConnect._tunnel + + def spy_tunnel(self, sock, context, deadline): + tunneled.append(sock) + original_tunnel(self, sock, context, deadline) + + with mock.patch.object(HTTPProxyKMSConnect, "_tunnel", spy_tunnel): + task = asyncio.create_task(AsyncHTTPProxyKMSConnect(host, port)(_kms_context())) + waited = await _run_blocking(requested.wait, 10) + self.assertTrue(waited, "proxy never received the CONNECT request") + task.cancel("no longer needed") + with self.assertRaises(asyncio.CancelledError): + await task + # Let the stub reply, completing the executor's future late. + reply.set() + await asyncio.sleep(0.5) + + self.assertEqual(len(tunneled), 1) + self.assertEqual(tunneled[0].fileno(), -1, "late socket was left open") + + async def test_connect_timeout_is_not_reclassified(self): + # A connect that times out keeps its socket.timeout type instead of + # being reported as a generic connect error. + def timeout_connect(self, address): + raise socket.timeout("timed out") + + with mock.patch.object(socket.socket, "connect", timeout_connect): + with self.assertRaises(socket.timeout): + HTTPProxyKMSConnect("127.0.0.1", 9999)._connect_proxy(time.monotonic() + 10) + + async def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): + # A proxy may coalesce its 200 with tunneled bytes; reading past the header would drop them. + def handler(conn): + conn.recv(4096) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") + + host, port = self._start_proxy(handler) + sock = await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) + self.addCleanup(sock.close) + sock.settimeout(10) + data = await _run_blocking(sock.recv, 64) + self.assertEqual(data, b"early-bytes") + + async def test_ipv6_host_is_bracketed_in_connect(self): + accepted: list[bytes] = [] + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + sock = await AsyncHTTPProxyKMSConnect(host, port)(_kms_context(host="::1")) + self.addCleanup(sock.close) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") + + async def test_oversized_connect_response_is_rejected(self): + def handler(conn): + conn.recv(4096) + # Never sends the terminator. + while True: + conn.sendall(b"x" * 1024) + + host, port = self._start_proxy(handler) + with self.assertRaisesRegex(OSError, "oversized CONNECT response"): + await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) + + async def test_remaining_raises_once_the_deadline_passes(self): + from pymongo.encryption_options import _remaining + + self.assertGreater(_remaining(time.monotonic() + 5), 0) + with self.assertRaises(socket.timeout): + _remaining(time.monotonic() - 1) + + async def test_bridge_failure_closes_the_proxy_socket(self): + # A failure inside _bridge must not strand the connected proxy socket. + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + tls.recv(4096) + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + tls.close() + + captured = [] + + def failing_bridge(self, proxy): + captured.append(proxy) + raise OSError("no file descriptors") + + host, port = self._start_proxy(handler) + context = _kms_context() + + with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): + with self.assertRaisesRegex(OSError, "no file descriptors"): + await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())(context) + + self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") + async def test_non_coroutine_callback_is_rejected(self): # A plain def must be rejected before it blocks the event loop. if _IS_SYNC: @@ -447,3 +796,183 @@ async def test_client_encryption_rejects_non_callable(self): OPTS, kms_connect_callback="not-callable", # type: ignore[arg-type] ) + + +class TestKmsConnectCallbackProse(AsyncEncryptionIntegrationTest): + @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") + async def asyncSetUp(self): + await super().asyncSetUp() + self.callback_calls: list[Any] = [] + + async def plain_callback(self, context): + self.callback_calls.append(context) + return await AsyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + def _proxy_tls_context(self): + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + # PYTHON-5040 tracks re-enabling verification once the test CA cert + # is fixed; the evergreen-tools CA lacks an Authority Key Identifier + # that newer OpenSSL requires, so verification fails on Windows 3.14. + ctx.verify_mode = ssl.CERT_NONE + return ctx + + async def tls_callback(self, context): + self.callback_calls.append(context) + callback = AsyncHTTPProxyKMSConnect( + KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, self._proxy_tls_context() + ) + return await callback(context) + + async def proxy_request(self, method, path, tls=False): + """Call the proxy's control endpoints and return the body.""" + if _IS_SYNC: + return self._proxy_request(method, path, tls) + return await asyncio.get_running_loop().run_in_executor( + None, self._proxy_request, method, path, tls + ) + + def _proxy_request(self, method, path, tls=False): + if tls: + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() + ) + else: + conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") + try: + conn.request(method, path) + return conn.getresponse().read().decode() + finally: + conn.close() + + async def connect_count(self, tls=False): + body = await self.proxy_request("GET", "/metrics", tls=tls) + # One "key value" per line; the server also emits connect_target. + for line in body.splitlines(): + key, _, value = line.partition(" ") + if key == "connect_count": + return int(value) + raise AssertionError(f"no connect_count in metrics body: {body!r}") + + async def test_01_plain_http_proxy(self): + await self.proxy_request("POST", "/reset") + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(await self.connect_count(), 1) + + async def test_02_https_proxy(self): + await self.proxy_request("POST", "/reset", tls=True) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.tls_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(await self.connect_count(tls=True), 1) + + async def test_03_auto_encryption_through_proxy(self): + await self.client.keyvault.datakeys.drop() + await self.client.db.coll.drop() + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + data_key_id = await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + schema = { + "bsonType": "object", + "properties": { + "encrypted_string": { + "encrypt": { + "keyId": [data_key_id], + "bsonType": "string", + "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", + } + } + }, + } + + await self.proxy_request("POST", "/reset") + opts = AutoEncryptionOpts( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + schema_map={"db.coll": schema}, + kms_connect_callback=self.plain_callback, + ) + client_encrypted = await self.async_rs_or_single_client(auto_encryption_opts=opts) + + await client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) + decrypted = await client_encrypted.db.coll.find_one({"_id": 1}) + self.assertEqual(decrypted["encrypted_string"], "hello") + + raw = await self.client.db.coll.find_one({"_id": 1}) + self.assertIsInstance(raw["encrypted_string"], Binary) + + # The decrypt reuses the cached key, so exactly one KMS request follows + # the reset. + self.assertEqual(await self.connect_count(), 1) + + async def test_04_callback_error(self): + async def failing_callback(context): + raise OSError("proxy is on fire") + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=failing_callback, + ) + with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + @unittest.skip( + "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " + "callback always receives the default KMS connect timeout" + ) + async def test_05_callback_receives_timeout(self): + key_vault_client = await self.async_rs_or_single_client(timeoutMS=1000) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + key_vault_client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + self.assertTrue(self.callback_calls, "callback was never invoked") + for context in self.callback_calls: + # Checks only the spec's non-zero requirement, which cannot fail. + self.assertIsNotNone(context.timeout) + self.assertGreater(context.timeout, 0) + + async def test_06_retry_after_network_error(self): + state = {"calls": 0} + + async def flaky_callback(context): + state["calls"] += 1 + if state["calls"] == 1: + raise OSError("first attempt fails") + return await AsyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=flaky_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(state["calls"], 2) diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index 8bd59ca689..fac9a49d43 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -1,9 +1,10 @@ -"""Tests for the KMS connect callback.""" +"""Tests for the KMS connect callback and HTTP proxy support.""" from __future__ import annotations import asyncio import dataclasses +import http.client import os import socket import ssl @@ -11,14 +12,18 @@ import time import unittest from asyncio.trsock import TransportSocket +from typing import Any from unittest import mock import pytest import pymongo +from bson.binary import Binary from pymongo.encryption_options import ( AutoEncryptionOpts, + HTTPProxyKMSConnect, KMSConnectContext, + SyncHTTPProxyKMSConnect, ) from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout from pymongo.pool_options import PoolOptions @@ -30,8 +35,8 @@ _wrap_encryption_errors, ) from test import PyMongoTestCase -from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM -from test.test_encryption import OPTS +from test.helpers_shared import AWS_CREDS, CA_PEM, CERT_PATH, CLIENT_PEM +from test.test_encryption import OPTS, EncryptionIntegrationTest _IS_SYNC = True @@ -39,6 +44,15 @@ _KMS_ADDRESS = ("kms.example.com", 443) +KMS_PROXY_HOST = "127.0.0.1" +KMS_PROXY_PORT = 9004 +KMS_TLS_PROXY_PORT = 9005 + +AWS_MASTER_KEY = { + "region": "us-east-1", + "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", +} + def _tls_server_context(cert=CLIENT_PEM): ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) @@ -69,6 +83,31 @@ def callback(context): return callback +def _kms_context(host="kms.example.com", port=443, timeout=10): + """A KMSConnectContext with the defaults used throughout these tests.""" + return KMSConnectContext(host=host, port=port, timeout=timeout) + + +def _read_http_request(conn): + """Read until the blank line that ends a CONNECT request, or None on EOF.""" + request = b"" + while b"\r\n\r\n" not in request: + chunk = conn.recv(4096) + if not chunk: + return None + request += chunk + return request + + +def _insecure_client_context(): + # PYTHON-5040 tracks re-enabling verification: the evergreen-tools CA + # lacks an Authority Key Identifier newer OpenSSL requires. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + return ctx + + class TestKmsConnectCallbackUnit(PyMongoTestCase): """Contract checks for kms_connect_callback that need no KMS server.""" @@ -89,6 +128,63 @@ def _socketpair(self): self.addCleanup(right.close) return left, right + def _start_proxy(self, handler, backlog=1): + """Serve each accepted connection with ``handler(conn)`` in a daemon thread.""" + listener = self._listen(backlog) + + def serve(): + for _ in range(backlog): + try: + conn, _ = listener.accept() + except OSError: + return + try: + handler(conn) + except OSError: + pass + finally: + conn.close() + + threading.Thread(target=serve, daemon=True).start() + return listener.getsockname() + + def _record_and_reply(self, accepted, reply): + """A proxy that records each CONNECT request, replies ``reply``, and closes.""" + + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + accepted.append(request) + conn.sendall(reply) + + return self._start_proxy(handler) + + def _tls_echo_proxy(self, delay=0): + """A TLS CONNECT proxy that replies 200, then echoes one tunneled read.""" + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + request = _read_http_request(tls) + if request is None: + return + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # The tunneled peer speaks only after the client does, as a TLS + # server would; ``delay`` lets the reply outlast the CONNECT deadline. + if delay: + time.sleep(delay) + tls.sendall(b"echo:" + tls.recv(64)) + tls.close() + + return self._start_proxy(handler) + + def _echo_over_tunnel(self, sock): + sock.settimeout(10) + _run_blocking(sock.sendall, b"ping") + data = _run_blocking(sock.recv, 64) + self.assertEqual(data, b"echo:ping") + def test_init_kms_connect_callback(self): opts = AutoEncryptionOpts({}, "k.d") self.assertIsNone(opts._kms_connect_callback) @@ -240,6 +336,257 @@ def test_cancelled_tls_wrap_closes_late_socket(self): _close_late_socket(future) self.assertEqual(left.fileno(), -1) + def test_http_proxy_helper_tunnels_and_reports_refusal(self): + # Covers the CONNECT handshake without KMS credentials. + accepted: list[bytes] = [] + context = _kms_context() + + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + sock = SyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n" + ) + with self.assertRaisesRegex(OSError, "refused CONNECT"): + SyncHTTPProxyKMSConnect(host, port)(context) + + # Any 2xx status is a successful tunnel, not just HTTP/1.1 200. + host, port = self._record_and_reply( + accepted, b"HTTP/1.0 200 Connection Established\r\n\r\n" + ) + sock = SyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + + # A status code must be exactly three digits, with no zero padding. + for reply in (b"HTTP/1.1 2000 Evil\r\n\r\n", b"HTTP/1.1 00200 Evil\r\n\r\n"): + host, port = self._record_and_reply(accepted, reply) + with self.assertRaisesRegex(OSError, "refused CONNECT"): + SyncHTTPProxyKMSConnect(host, port)(context) + + def test_control_characters_in_kms_host_are_rejected(self): + # Reject CR/LF in the configurable host before it reaches CONNECT. + callback = SyncHTTPProxyKMSConnect("proxy.example.com", 8080) + context = _kms_context(host="kms.example.com\r\nX-Injected: 1") + with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): + callback(context) + # Whitespace would split the request line into extra tokens. + context = _kms_context(host="kms.example.com ") + with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): + callback(context) + + def test_http_proxy_helper_sends_custom_headers(self): + # Extra CONNECT headers reach the proxy verbatim. + accepted: list[bytes] = [] + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Trace-Id": "abc123"} + sock = SyncHTTPProxyKMSConnect(host, port, headers=headers)(_kms_context()) + self.addCleanup(sock.close) + request = accepted[0] + self.assertEqual(request.split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + self.assertIn(b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n", request) + self.assertIn(b"\r\nX-Trace-Id: abc123\r\n", request) + self.assertEqual(request.count(b"\r\nHost: "), 1) + + def test_http_proxy_helper_authenticates_to_the_proxy(self): + # The motivating case: 407 without credentials, 200 with them. + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + if b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request: + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + else: + conn.sendall(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") + + host, port = self._start_proxy(handler, backlog=2) + context = _kms_context() + with self.assertRaisesRegex(OSError, "refused CONNECT"): + SyncHTTPProxyKMSConnect(host, port)(context) + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz"} + sock = SyncHTTPProxyKMSConnect(host, port, headers=headers)(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + + def test_http_proxy_helper_rejects_bad_headers(self): + for headers in [ + {"Bad\r\nName": "x"}, + {"Bad Name": "x"}, + {"Bad\tName": "x"}, + {"X-Ok": "ok\r\nInjected: 1"}, + {"Host": "evil.example.com"}, + {"host": "evil.example.com"}, + {"": "x"}, + {"Bad:Name": "x"}, + ]: + with self.assertRaisesRegex(ConfigurationError, "proxy header|Host CONNECT header"): + SyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) + + for headers in [{1: "x"}, {"X-Ok": 1}, {None: "x"}, {"X-Ok": None}]: + with self.assertRaisesRegex(TypeError, "must be strings"): + SyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) + + def test_http_proxy_helper_accepts_legal_header_values(self): + # Colons and spaces are legal in values (e.g. auth schemes); only + # CR/LF would let a value inject a request line. + callback = SyncHTTPProxyKMSConnect( + "proxy.example.com", + 8080, + headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, + ) + self.assertEqual( + callback.headers, + {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, + ) + + def test_tls_proxy_helper_bridges_the_tunnel(self): + # Covers the TLS-proxy path and the socketpair relay without KMS creds. + host, port = self._tls_echo_proxy() + sock = HTTPProxyKMSConnect(host, port, _insecure_client_context())(_kms_context()) + self.addCleanup(sock.close) + self._echo_over_tunnel(sock) + + def test_bridge_does_not_inherit_the_connect_deadline(self): + # The relay must outlast the much shorter CONNECT deadline. + host, port = self._tls_echo_proxy(delay=3.0) + sock = HTTPProxyKMSConnect(host, port, _insecure_client_context())( + _kms_context(timeout=2.0) + ) + self.addCleanup(sock.close) + self._echo_over_tunnel(sock) + + def test_proxy_closing_before_connect_reply_raises(self): + def handler(conn): + # Read the CONNECT request, then hang up without replying. + conn.recv(4096) + + host, port = self._start_proxy(handler) + with self.assertRaisesRegex(OSError, "proxy closed the connection"): + SyncHTTPProxyKMSConnect(host, port)(_kms_context()) + + def test_cancelled_proxy_connect_closes_the_late_socket(self): + # A cancelled connect must close the socket the executor thread + # produces after the cancellation. + if _IS_SYNC: + raise unittest.SkipTest("cancellation is an async-only behavior") + + requested = threading.Event() + reply = threading.Event() + + def handler(conn): + conn.recv(4096) + requested.set() + if not reply.wait(10): + return + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # Keep the connection open so the tunnel can complete its reads. + time.sleep(0.1) + + host, port = self._start_proxy(handler) + + tunneled: list[socket.socket] = [] + original_tunnel = HTTPProxyKMSConnect._tunnel + + def spy_tunnel(self, sock, context, deadline): + tunneled.append(sock) + original_tunnel(self, sock, context, deadline) + + with mock.patch.object(HTTPProxyKMSConnect, "_tunnel", spy_tunnel): + task = asyncio.create_task(SyncHTTPProxyKMSConnect(host, port)(_kms_context())) + waited = _run_blocking(requested.wait, 10) + self.assertTrue(waited, "proxy never received the CONNECT request") + task.cancel("no longer needed") + with self.assertRaises(asyncio.CancelledError): + task + # Let the stub reply, completing the executor's future late. + reply.set() + time.sleep(0.5) + + self.assertEqual(len(tunneled), 1) + self.assertEqual(tunneled[0].fileno(), -1, "late socket was left open") + + def test_connect_timeout_is_not_reclassified(self): + # A connect that times out keeps its socket.timeout type instead of + # being reported as a generic connect error. + def timeout_connect(self, address): + raise socket.timeout("timed out") + + with mock.patch.object(socket.socket, "connect", timeout_connect): + with self.assertRaises(socket.timeout): + HTTPProxyKMSConnect("127.0.0.1", 9999)._connect_proxy(time.monotonic() + 10) + + def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): + # A proxy may coalesce its 200 with tunneled bytes; reading past the header would drop them. + def handler(conn): + conn.recv(4096) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") + + host, port = self._start_proxy(handler) + sock = SyncHTTPProxyKMSConnect(host, port)(_kms_context()) + self.addCleanup(sock.close) + sock.settimeout(10) + data = _run_blocking(sock.recv, 64) + self.assertEqual(data, b"early-bytes") + + def test_ipv6_host_is_bracketed_in_connect(self): + accepted: list[bytes] = [] + host, port = self._record_and_reply( + accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" + ) + sock = SyncHTTPProxyKMSConnect(host, port)(_kms_context(host="::1")) + self.addCleanup(sock.close) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") + + def test_oversized_connect_response_is_rejected(self): + def handler(conn): + conn.recv(4096) + # Never sends the terminator. + while True: + conn.sendall(b"x" * 1024) + + host, port = self._start_proxy(handler) + with self.assertRaisesRegex(OSError, "oversized CONNECT response"): + SyncHTTPProxyKMSConnect(host, port)(_kms_context()) + + def test_remaining_raises_once_the_deadline_passes(self): + from pymongo.encryption_options import _remaining + + self.assertGreater(_remaining(time.monotonic() + 5), 0) + with self.assertRaises(socket.timeout): + _remaining(time.monotonic() - 1) + + def test_bridge_failure_closes_the_proxy_socket(self): + # A failure inside _bridge must not strand the connected proxy socket. + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + tls.recv(4096) + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + tls.close() + + captured = [] + + def failing_bridge(self, proxy): + captured.append(proxy) + raise OSError("no file descriptors") + + host, port = self._start_proxy(handler) + context = _kms_context() + + with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): + with self.assertRaisesRegex(OSError, "no file descriptors"): + HTTPProxyKMSConnect(host, port, _insecure_client_context())(context) + + self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") + def test_non_coroutine_callback_is_rejected(self): # A plain def must be rejected before it blocks the event loop. if _IS_SYNC: @@ -445,3 +792,183 @@ def test_client_encryption_rejects_non_callable(self): OPTS, kms_connect_callback="not-callable", # type: ignore[arg-type] ) + + +class TestKmsConnectCallbackProse(EncryptionIntegrationTest): + @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") + def setUp(self): + super().setUp() + self.callback_calls: list[Any] = [] + + def plain_callback(self, context): + self.callback_calls.append(context) + return SyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + def _proxy_tls_context(self): + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + # PYTHON-5040 tracks re-enabling verification once the test CA cert + # is fixed; the evergreen-tools CA lacks an Authority Key Identifier + # that newer OpenSSL requires, so verification fails on Windows 3.14. + ctx.verify_mode = ssl.CERT_NONE + return ctx + + def tls_callback(self, context): + self.callback_calls.append(context) + callback = SyncHTTPProxyKMSConnect( + KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, self._proxy_tls_context() + ) + return callback(context) + + def proxy_request(self, method, path, tls=False): + """Call the proxy's control endpoints and return the body.""" + if _IS_SYNC: + return self._proxy_request(method, path, tls) + return asyncio.get_running_loop().run_in_executor( + None, self._proxy_request, method, path, tls + ) + + def _proxy_request(self, method, path, tls=False): + if tls: + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() + ) + else: + conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") + try: + conn.request(method, path) + return conn.getresponse().read().decode() + finally: + conn.close() + + def connect_count(self, tls=False): + body = self.proxy_request("GET", "/metrics", tls=tls) + # One "key value" per line; the server also emits connect_target. + for line in body.splitlines(): + key, _, value = line.partition(" ") + if key == "connect_count": + return int(value) + raise AssertionError(f"no connect_count in metrics body: {body!r}") + + def test_01_plain_http_proxy(self): + self.proxy_request("POST", "/reset") + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(self.connect_count(), 1) + + def test_02_https_proxy(self): + self.proxy_request("POST", "/reset", tls=True) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.tls_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(self.connect_count(tls=True), 1) + + def test_03_auto_encryption_through_proxy(self): + self.client.keyvault.datakeys.drop() + self.client.db.coll.drop() + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + data_key_id = encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + schema = { + "bsonType": "object", + "properties": { + "encrypted_string": { + "encrypt": { + "keyId": [data_key_id], + "bsonType": "string", + "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", + } + } + }, + } + + self.proxy_request("POST", "/reset") + opts = AutoEncryptionOpts( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + schema_map={"db.coll": schema}, + kms_connect_callback=self.plain_callback, + ) + client_encrypted = self.rs_or_single_client(auto_encryption_opts=opts) + + client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) + decrypted = client_encrypted.db.coll.find_one({"_id": 1}) + self.assertEqual(decrypted["encrypted_string"], "hello") + + raw = self.client.db.coll.find_one({"_id": 1}) + self.assertIsInstance(raw["encrypted_string"], Binary) + + # The decrypt reuses the cached key, so exactly one KMS request follows + # the reset. + self.assertEqual(self.connect_count(), 1) + + def test_04_callback_error(self): + def failing_callback(context): + raise OSError("proxy is on fire") + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=failing_callback, + ) + with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + @unittest.skip( + "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " + "callback always receives the default KMS connect timeout" + ) + def test_05_callback_receives_timeout(self): + key_vault_client = self.rs_or_single_client(timeoutMS=1000) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + key_vault_client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + self.assertTrue(self.callback_calls, "callback was never invoked") + for context in self.callback_calls: + # Checks only the spec's non-zero requirement, which cannot fail. + self.assertIsNotNone(context.timeout) + self.assertGreater(context.timeout, 0) + + def test_06_retry_after_network_error(self): + state = {"calls": 0} + + def flaky_callback(context): + state["calls"] += 1 + if state["calls"] == 1: + raise OSError("first attempt fails") + return SyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=flaky_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(state["calls"], 2) From 1a4ccb34dbfe99ee3c69c72587833ab3b1acbf86 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Thu, 8 Oct 2026 17:08:47 -0500 Subject: [PATCH 2/4] PYTHON-6154 Fix the KMS connect refactor Deliver the PR's stated goals: - Rename pymongo/_kms_connect.py to pymongo/_kms_connect_shared.py, so the shared module reads naturally next to the per-flavor pymongo/{asynchronous,synchronous}/_kms_connect.py modules. - Genericize the shared module docstring: it enumerated the module's exact contents, which would go stale as helpers move in. - Drop the PYTHON-6147 feature tests (HTTP proxy helpers and the prose class) that leaked into this refactor and broke test collection: they import HTTPProxyKMSConnect/AsyncHTTPProxyKMSConnect, which are added by PYTHON-6147, not by this PR. - Replace the synchro-mirrored test pair with a single hand-written test/test_kms_connect.py, written once and parameterized over both APIs through a Flavor facade (the 18 pre-existing tests, 32 variants, byte-identical bodies). The async-side file is gone, so synchro no longer mirrors it. --- ..._kms_connect.py => _kms_connect_shared.py} | 8 +- pymongo/asynchronous/_kms_connect.py | 4 +- pymongo/asynchronous/encryption.py | 2 +- pymongo/encryption_options.py | 2 +- pymongo/synchronous/_kms_connect.py | 4 +- pymongo/synchronous/encryption.py | 2 +- test/asynchronous/test_kms_connect.py | 978 ------------- test/test_kms_connect.py | 1228 ++++++----------- 8 files changed, 440 insertions(+), 1788 deletions(-) rename pymongo/{_kms_connect.py => _kms_connect_shared.py} (88%) delete mode 100644 test/asynchronous/test_kms_connect.py diff --git a/pymongo/_kms_connect.py b/pymongo/_kms_connect_shared.py similarity index 88% rename from pymongo/_kms_connect.py rename to pymongo/_kms_connect_shared.py index 50cd3d73f0..710f380d92 100644 --- a/pymongo/_kms_connect.py +++ b/pymongo/_kms_connect_shared.py @@ -12,12 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""KMS connection helpers, shared by both the synchronous and asynchronous APIs. +"""KMS connection support shared by the synchronous and asynchronous APIs. -Holds the ``KMSConnectContext`` passed to ``kms_connect_callback`` and the -callback type aliases. The per-API ``_connect_kms`` that connects to the KMS -host and performs the TLS handshake lives in -``pymongo.asynchronous._kms_connect`` and its generated synchronous mirror. +The per-API connection logic lives in ``pymongo.asynchronous._kms_connect`` +and its generated synchronous mirror. """ from __future__ import annotations diff --git a/pymongo/asynchronous/_kms_connect.py b/pymongo/asynchronous/_kms_connect.py index 176710e67d..3f4ba1a7d3 100644 --- a/pymongo/asynchronous/_kms_connect.py +++ b/pymongo/asynchronous/_kms_connect.py @@ -17,7 +17,7 @@ Connects to a KMS host, directly or through a ``kms_connect_callback``, and performs the KMS TLS handshake over the connection. The synchronous mirror of this module is generated by synchro. The helpers shared by both APIs live in -``pymongo._kms_connect``. +``pymongo._kms_connect_shared``. """ from __future__ import annotations @@ -29,7 +29,7 @@ from typing import TYPE_CHECKING, Any, Optional, Union, cast from pymongo import _csot -from pymongo._kms_connect import ( +from pymongo._kms_connect_shared import ( AsyncKMSConnectCallback, KMSConnectContext, _close_rejected_kms_socket, diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 955529fa97..9c2d72604c 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -53,7 +53,7 @@ from bson.codec_options import CodecOptions from bson.raw_bson import DEFAULT_RAW_BSON_OPTIONS, RawBSONDocument, _inflate_bson from pymongo import _csot, _op_id -from pymongo._kms_connect import AsyncKMSConnectCallback +from pymongo._kms_connect_shared import AsyncKMSConnectCallback from pymongo.asynchronous._kms_connect import _KMS_CONNECT_TIMEOUT, _connect_kms from pymongo.asynchronous.collection import AsyncCollection from pymongo.asynchronous.cursor import AsyncCursor diff --git a/pymongo/encryption_options.py b/pymongo/encryption_options.py index 41e200eefc..852ce37210 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -35,7 +35,7 @@ except ImportError: _HAVE_PYMONGOCRYPT = False from bson import int64 -from pymongo._kms_connect import ( # noqa: F401 +from pymongo._kms_connect_shared import ( # noqa: F401 AsyncKMSConnectCallback, KMSConnectCallback, KMSConnectContext, diff --git a/pymongo/synchronous/_kms_connect.py b/pymongo/synchronous/_kms_connect.py index bd453e9114..b05ba019a4 100644 --- a/pymongo/synchronous/_kms_connect.py +++ b/pymongo/synchronous/_kms_connect.py @@ -17,7 +17,7 @@ Connects to a KMS host, directly or through a ``kms_connect_callback``, and performs the KMS TLS handshake over the connection. The synchronous mirror of this module is generated by synchro. The helpers shared by both APIs live in -``pymongo._kms_connect``. +``pymongo._kms_connect_shared``. """ from __future__ import annotations @@ -29,7 +29,7 @@ from typing import TYPE_CHECKING, Any, Optional, Union, cast from pymongo import _csot -from pymongo._kms_connect import ( +from pymongo._kms_connect_shared import ( KMSConnectCallback, KMSConnectContext, _close_rejected_kms_socket, diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index 9d0b12c683..d73aec4b68 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -52,7 +52,7 @@ from bson.codec_options import CodecOptions from bson.raw_bson import DEFAULT_RAW_BSON_OPTIONS, RawBSONDocument, _inflate_bson from pymongo import _csot, _op_id -from pymongo._kms_connect import KMSConnectCallback +from pymongo._kms_connect_shared import KMSConnectCallback from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( AutoEncryptionOpts, diff --git a/test/asynchronous/test_kms_connect.py b/test/asynchronous/test_kms_connect.py deleted file mode 100644 index a1a6ee2205..0000000000 --- a/test/asynchronous/test_kms_connect.py +++ /dev/null @@ -1,978 +0,0 @@ -"""Tests for the KMS connect callback and HTTP proxy support.""" - -from __future__ import annotations - -import asyncio -import dataclasses -import http.client -import os -import socket -import ssl -import threading -import time -import unittest -from asyncio.trsock import TransportSocket -from typing import Any -from unittest import mock - -import pytest - -import pymongo -from bson.binary import Binary -from pymongo.asynchronous.encryption import ( - AsyncClientEncryption, - _connect_kms, - _EncryptionIO, - _wrap_encryption_errors, -) -from pymongo.encryption_options import ( - AsyncHTTPProxyKMSConnect, - AutoEncryptionOpts, - HTTPProxyKMSConnect, - KMSConnectContext, -) -from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout -from pymongo.pool_options import PoolOptions -from pymongo.ssl_support import get_ssl_context -from test.asynchronous import AsyncPyMongoTestCase -from test.asynchronous.test_encryption import OPTS, AsyncEncryptionIntegrationTest -from test.helpers_shared import AWS_CREDS, CA_PEM, CERT_PATH, CLIENT_PEM - -_IS_SYNC = False - -pytestmark = pytest.mark.encryption - -_KMS_ADDRESS = ("kms.example.com", 443) - -KMS_PROXY_HOST = "127.0.0.1" -KMS_PROXY_PORT = 9004 -KMS_TLS_PROXY_PORT = 9005 - -AWS_MASTER_KEY = { - "region": "us-east-1", - "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", -} - - -def _tls_server_context(cert=CLIENT_PEM): - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(cert) - return ctx - - -def _client_tls_context(verify=False): - # verify=False matches the driver's test mode: the local certs don't verify. - if verify: - return get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) - return get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) - - -async def _run_blocking(func, *args): - """Run a blocking callable off the event loop (inline in the sync version).""" - if _IS_SYNC: - return func(*args) - return await asyncio.get_running_loop().run_in_executor(None, func, *args) - - -def _callback_returning(value): - """A kms_connect_callback that always produces ``value``.""" - - async def callback(context): - return value - - return callback - - -def _kms_context(host="kms.example.com", port=443, timeout=10): - """A KMSConnectContext with the defaults used throughout these tests.""" - return KMSConnectContext(host=host, port=port, timeout=timeout) - - -def _read_http_request(conn): - """Read until the blank line that ends a CONNECT request, or None on EOF.""" - request = b"" - while b"\r\n\r\n" not in request: - chunk = conn.recv(4096) - if not chunk: - return None - request += chunk - return request - - -def _insecure_client_context(): - # PYTHON-5040 tracks re-enabling verification: the evergreen-tools CA - # lacks an Authority Key Identifier newer OpenSSL requires. - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ctx.check_hostname = False - ctx.verify_mode = ssl.CERT_NONE - return ctx - - -class TestKmsConnectCallbackUnit(AsyncPyMongoTestCase): - """Contract checks for kms_connect_callback that need no KMS server.""" - - @staticmethod - def _pool_options(ssl_context=None): - return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=ssl_context) - - def _listen(self, backlog=1): - listener = socket.socket() - listener.bind(("127.0.0.1", 0)) - listener.listen(backlog) - self.addCleanup(listener.close) - return listener - - def _socketpair(self): - left, right = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) - return left, right - - def _start_proxy(self, handler, backlog=1): - """Serve each accepted connection with ``handler(conn)`` in a daemon thread.""" - listener = self._listen(backlog) - - def serve(): - for _ in range(backlog): - try: - conn, _ = listener.accept() - except OSError: - return - try: - handler(conn) - except OSError: - pass - finally: - conn.close() - - threading.Thread(target=serve, daemon=True).start() - return listener.getsockname() - - def _record_and_reply(self, accepted, reply): - """A proxy that records each CONNECT request, replies ``reply``, and closes.""" - - def handler(conn): - request = _read_http_request(conn) - if request is None: - return - accepted.append(request) - conn.sendall(reply) - - return self._start_proxy(handler) - - def _tls_echo_proxy(self, delay=0): - """A TLS CONNECT proxy that replies 200, then echoes one tunneled read.""" - server_ctx = _tls_server_context() - - def handler(conn): - tls = server_ctx.wrap_socket(conn, server_side=True) - request = _read_http_request(tls) - if request is None: - return - tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - # The tunneled peer speaks only after the client does, as a TLS - # server would; ``delay`` lets the reply outlast the CONNECT deadline. - if delay: - time.sleep(delay) - tls.sendall(b"echo:" + tls.recv(64)) - tls.close() - - return self._start_proxy(handler) - - async def _echo_over_tunnel(self, sock): - sock.settimeout(10) - await _run_blocking(sock.sendall, b"ping") - data = await _run_blocking(sock.recv, 64) - self.assertEqual(data, b"echo:ping") - - async def test_init_kms_connect_callback(self): - opts = AutoEncryptionOpts({}, "k.d") - self.assertIsNone(opts._kms_connect_callback) - - async def callback(context): - raise AssertionError("not called") - - opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) - self.assertIs(opts._kms_connect_callback, callback) - - for bad in [1, "not-callable", object()]: - with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): - AutoEncryptionOpts({}, "k.d", kms_connect_callback=bad) # type: ignore[arg-type] - - context = KMSConnectContext(host="kms.example.com", port=443, timeout=9.5) - self.assertEqual(context.host, "kms.example.com") - self.assertEqual(context.port, 443) - self.assertEqual(context.timeout, 9.5) - with self.assertRaises(dataclasses.FrozenInstanceError): - context.host = "evil.example.com" # type: ignore[misc] - - async def test_non_socket_return_raises_configuration_error(self): - with self.assertRaisesRegex(ConfigurationError, "must return a connected"): - await _connect_kms( - _KMS_ADDRESS, self._pool_options(), _callback_returning("not-a-socket"), 10.0 - ) - - async def test_already_wrapped_socket_is_rejected(self): - # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ctx.check_hostname = False - ctx.verify_mode = ssl.CERT_NONE - left, _right = self._socketpair() - # No peer needed to produce a genuine ssl.SSLSocket. - wrapped = ctx.wrap_socket(left, do_handshake_on_connect=False, server_hostname="x") - self.addCleanup(wrapped.close) - - with self.assertRaisesRegex(ConfigurationError, "unwrapped"): - await _connect_kms( - _KMS_ADDRESS, self._pool_options(), _callback_returning(wrapped), 10.0 - ) - - async def test_context_receives_host_port_and_timeout(self): - received = [] - left, _right = self._socketpair() - - async def callback(context): - received.append(context) - return left - - # ssl_context=None returns the socket unchanged, so a plain socket is accepted. - conn = await _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 12.5) - self.assertIs(conn, left) - - self.assertEqual(len(received), 1) - self.assertEqual(received[0].host, "kms.example.com") - self.assertEqual(received[0].port, 443) - self.assertEqual(received[0].timeout, 12.5) - - async def test_non_blocking_socket_from_callback_is_accepted(self): - # Without the driver normalizing the mode, this raises ValueError. - server_ctx = _tls_server_context() - listener = self._listen() - - def serve(): - try: - conn, _ = listener.accept() - server_ctx.wrap_socket(conn, server_side=True).close() - except OSError: - pass - - threading.Thread(target=serve, daemon=True).start() - - options = self._pool_options(_client_tls_context()) - - def connect(): - sock = socket.create_connection(listener.getsockname(), timeout=10) - sock.setblocking(False) - return sock - - async def callback(context): - return await _run_blocking(connect) - - conn = await _connect_kms(listener.getsockname(), options, callback, 10.0) - self.addCleanup(conn.close) - self.assertIsNotNone(conn.gettimeout()) - - async def test_tls_verification_targets_the_kms_host(self): - # The handshake must verify against the KMS address, not the peer the - # callback connected to. The server cert covers 127.0.0.1 (the peer) - # and localhost, but not the KMS hostname used below, so only - # address-based verification produces this outcome. - server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) - listener = self._listen(2) - - def serve(): - for _ in range(2): - try: - conn, _ = listener.accept() - server_ctx.wrap_socket(conn, server_side=True).close() - except (OSError, ssl.SSLError): - # The mismatched-name attempt fails mid-handshake. - pass - - threading.Thread(target=serve, daemon=True).start() - - # Full verification: trusted CA, invalid certs and hostnames rejected. - options = self._pool_options(_client_tls_context(verify=True)) - - created = [] - - def connect(): - sock = socket.create_connection(listener.getsockname(), timeout=10) - created.append(sock) - return sock - - async def callback(context): - return await _run_blocking(connect) - - port = listener.getsockname()[1] - # The cert covers localhost: verifying against the KMS address succeeds. - conn = await _connect_kms(("localhost", port), options, callback, 10.0) - self.addCleanup(conn.close) - # TLS-wrapped in either SSL flavor: a new object, not the plain socket. - self.assertIsNot(conn, created[0]) - # The cert does not cover this name: verification must fail even - # though the peer (127.0.0.1) presents a cert valid for itself. - with self.assertRaises(ConnectionFailure): - await _connect_kms(("kms.example.com", port), options, callback, 10.0) - - async def test_asyncio_transport_socket_is_rejected(self): - # get_extra_info("socket") is a TransportSocket, not a socket.socket. - left, _right = self._socketpair() - - with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): - await _connect_kms( - _KMS_ADDRESS, self._pool_options(), _callback_returning(TransportSocket(left)), 10.0 - ) - - async def test_cancelled_tls_wrap_closes_late_socket(self): - # A cancelled wrap can leave the executor producing an SSLSocket; the - # done callback must close it. - if _IS_SYNC: - raise unittest.SkipTest("the cancel-safe wrap is an async path") - from pymongo.pool_shared import _close_late_socket - - left, _right = self._socketpair() - future = asyncio.get_running_loop().create_future() - future.set_result(left) - self.assertNotEqual(left.fileno(), -1) - _close_late_socket(future) - self.assertEqual(left.fileno(), -1) - - async def test_http_proxy_helper_tunnels_and_reports_refusal(self): - # Covers the CONNECT handshake without KMS credentials. - accepted: list[bytes] = [] - context = _kms_context() - - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - sock = await AsyncHTTPProxyKMSConnect(host, port)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") - - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n" - ) - with self.assertRaisesRegex(OSError, "refused CONNECT"): - await AsyncHTTPProxyKMSConnect(host, port)(context) - - # Any 2xx status is a successful tunnel, not just HTTP/1.1 200. - host, port = self._record_and_reply( - accepted, b"HTTP/1.0 200 Connection Established\r\n\r\n" - ) - sock = await AsyncHTTPProxyKMSConnect(host, port)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - - # A status code must be exactly three digits, with no zero padding. - for reply in (b"HTTP/1.1 2000 Evil\r\n\r\n", b"HTTP/1.1 00200 Evil\r\n\r\n"): - host, port = self._record_and_reply(accepted, reply) - with self.assertRaisesRegex(OSError, "refused CONNECT"): - await AsyncHTTPProxyKMSConnect(host, port)(context) - - async def test_control_characters_in_kms_host_are_rejected(self): - # Reject CR/LF in the configurable host before it reaches CONNECT. - callback = AsyncHTTPProxyKMSConnect("proxy.example.com", 8080) - context = _kms_context(host="kms.example.com\r\nX-Injected: 1") - with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): - await callback(context) - # Whitespace would split the request line into extra tokens. - context = _kms_context(host="kms.example.com ") - with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): - await callback(context) - - async def test_http_proxy_helper_sends_custom_headers(self): - # Extra CONNECT headers reach the proxy verbatim. - accepted: list[bytes] = [] - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Trace-Id": "abc123"} - sock = await AsyncHTTPProxyKMSConnect(host, port, headers=headers)(_kms_context()) - self.addCleanup(sock.close) - request = accepted[0] - self.assertEqual(request.split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") - self.assertIn(b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n", request) - self.assertIn(b"\r\nX-Trace-Id: abc123\r\n", request) - self.assertEqual(request.count(b"\r\nHost: "), 1) - - async def test_http_proxy_helper_authenticates_to_the_proxy(self): - # The motivating case: 407 without credentials, 200 with them. - def handler(conn): - request = _read_http_request(conn) - if request is None: - return - if b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request: - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - else: - conn.sendall(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") - - host, port = self._start_proxy(handler, backlog=2) - context = _kms_context() - with self.assertRaisesRegex(OSError, "refused CONNECT"): - await AsyncHTTPProxyKMSConnect(host, port)(context) - headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz"} - sock = await AsyncHTTPProxyKMSConnect(host, port, headers=headers)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - - async def test_http_proxy_helper_rejects_bad_headers(self): - for headers in [ - {"Bad\r\nName": "x"}, - {"Bad Name": "x"}, - {"Bad\tName": "x"}, - {"X-Ok": "ok\r\nInjected: 1"}, - {"Host": "evil.example.com"}, - {"host": "evil.example.com"}, - {"": "x"}, - {"Bad:Name": "x"}, - ]: - with self.assertRaisesRegex(ConfigurationError, "proxy header|Host CONNECT header"): - AsyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) - - for headers in [{1: "x"}, {"X-Ok": 1}, {None: "x"}, {"X-Ok": None}]: - with self.assertRaisesRegex(TypeError, "must be strings"): - AsyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) - - async def test_http_proxy_helper_accepts_legal_header_values(self): - # Colons and spaces are legal in values (e.g. auth schemes); only - # CR/LF would let a value inject a request line. - callback = AsyncHTTPProxyKMSConnect( - "proxy.example.com", - 8080, - headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, - ) - self.assertEqual( - callback.headers, - {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, - ) - - async def test_tls_proxy_helper_bridges_the_tunnel(self): - # Covers the TLS-proxy path and the socketpair relay without KMS creds. - host, port = self._tls_echo_proxy() - sock = await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())( - _kms_context() - ) - self.addCleanup(sock.close) - await self._echo_over_tunnel(sock) - - async def test_bridge_does_not_inherit_the_connect_deadline(self): - # The relay must outlast the much shorter CONNECT deadline. - host, port = self._tls_echo_proxy(delay=3.0) - sock = await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())( - _kms_context(timeout=2.0) - ) - self.addCleanup(sock.close) - await self._echo_over_tunnel(sock) - - async def test_proxy_closing_before_connect_reply_raises(self): - def handler(conn): - # Read the CONNECT request, then hang up without replying. - conn.recv(4096) - - host, port = self._start_proxy(handler) - with self.assertRaisesRegex(OSError, "proxy closed the connection"): - await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) - - async def test_cancelled_proxy_connect_closes_the_late_socket(self): - # A cancelled connect must close the socket the executor thread - # produces after the cancellation. - if _IS_SYNC: - raise unittest.SkipTest("cancellation is an async-only behavior") - - requested = threading.Event() - reply = threading.Event() - - def handler(conn): - conn.recv(4096) - requested.set() - if not reply.wait(10): - return - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - # Keep the connection open so the tunnel can complete its reads. - time.sleep(0.1) - - host, port = self._start_proxy(handler) - - tunneled: list[socket.socket] = [] - original_tunnel = HTTPProxyKMSConnect._tunnel - - def spy_tunnel(self, sock, context, deadline): - tunneled.append(sock) - original_tunnel(self, sock, context, deadline) - - with mock.patch.object(HTTPProxyKMSConnect, "_tunnel", spy_tunnel): - task = asyncio.create_task(AsyncHTTPProxyKMSConnect(host, port)(_kms_context())) - waited = await _run_blocking(requested.wait, 10) - self.assertTrue(waited, "proxy never received the CONNECT request") - task.cancel("no longer needed") - with self.assertRaises(asyncio.CancelledError): - await task - # Let the stub reply, completing the executor's future late. - reply.set() - await asyncio.sleep(0.5) - - self.assertEqual(len(tunneled), 1) - self.assertEqual(tunneled[0].fileno(), -1, "late socket was left open") - - async def test_connect_timeout_is_not_reclassified(self): - # A connect that times out keeps its socket.timeout type instead of - # being reported as a generic connect error. - def timeout_connect(self, address): - raise socket.timeout("timed out") - - with mock.patch.object(socket.socket, "connect", timeout_connect): - with self.assertRaises(socket.timeout): - HTTPProxyKMSConnect("127.0.0.1", 9999)._connect_proxy(time.monotonic() + 10) - - async def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): - # A proxy may coalesce its 200 with tunneled bytes; reading past the header would drop them. - def handler(conn): - conn.recv(4096) - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") - - host, port = self._start_proxy(handler) - sock = await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) - self.addCleanup(sock.close) - sock.settimeout(10) - data = await _run_blocking(sock.recv, 64) - self.assertEqual(data, b"early-bytes") - - async def test_ipv6_host_is_bracketed_in_connect(self): - accepted: list[bytes] = [] - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - sock = await AsyncHTTPProxyKMSConnect(host, port)(_kms_context(host="::1")) - self.addCleanup(sock.close) - self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") - - async def test_oversized_connect_response_is_rejected(self): - def handler(conn): - conn.recv(4096) - # Never sends the terminator. - while True: - conn.sendall(b"x" * 1024) - - host, port = self._start_proxy(handler) - with self.assertRaisesRegex(OSError, "oversized CONNECT response"): - await AsyncHTTPProxyKMSConnect(host, port)(_kms_context()) - - async def test_remaining_raises_once_the_deadline_passes(self): - from pymongo.encryption_options import _remaining - - self.assertGreater(_remaining(time.monotonic() + 5), 0) - with self.assertRaises(socket.timeout): - _remaining(time.monotonic() - 1) - - async def test_bridge_failure_closes_the_proxy_socket(self): - # A failure inside _bridge must not strand the connected proxy socket. - server_ctx = _tls_server_context() - - def handler(conn): - tls = server_ctx.wrap_socket(conn, server_side=True) - tls.recv(4096) - tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - tls.close() - - captured = [] - - def failing_bridge(self, proxy): - captured.append(proxy) - raise OSError("no file descriptors") - - host, port = self._start_proxy(handler) - context = _kms_context() - - with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): - with self.assertRaisesRegex(OSError, "no file descriptors"): - await AsyncHTTPProxyKMSConnect(host, port, _insecure_client_context())(context) - - self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") - - async def test_non_coroutine_callback_is_rejected(self): - # A plain def must be rejected before it blocks the event loop. - if _IS_SYNC: - raise unittest.SkipTest("a regular function is correct for the sync API") - - entered = [] - - def callback(context): - entered.append(context) - return None - - with self.assertRaisesRegex(ConfigurationError, "coroutine function"): - await _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) - self.assertEqual(entered, [], "invalid callback must not be entered") - - async def test_unconnected_socket_from_callback_is_rejected(self): - # An unconnected socket would fail later as a transient error and be retried. - bare = socket.socket() - self.addCleanup(bare.close) - - with self.assertRaisesRegex(ConfigurationError, "already connected"): - await _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(bare), 10.0) - - async def test_datagram_socket_from_callback_is_rejected(self): - # TLS on a connected UDP socket raises NotImplementedError, which would be retried. - left = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - right = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - self.addCleanup(left.close) - self.addCleanup(right.close) - right.bind(("127.0.0.1", 0)) - left.connect(right.getsockname()) - - with self.assertRaisesRegex(ConfigurationError, "stream socket"): - await _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(left), 10.0) - - async def test_kms_request_does_not_retry_a_contract_violation(self): - # _connect_kms has no retry loop; the no-retry guarantee is in - # kms_request, so exercise that instead. - calls = [] - - async def callback(context): - calls.append(context) - return "not-a-socket" - - opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) - io = _EncryptionIO(None, mock.MagicMock(), None, opts) - - class StubKmsContext: - endpoint = "kms.example.com:443" - message = b"request" - kms_provider = "aws" - usleep = 0 - bytes_needed = 1 - - def feed(self, data): - raise AssertionError("should not reach the socket") - - def fail(self): - raise AssertionError("a contract violation must not be retried") - - with self.assertRaises(ConfigurationError): - await io.kms_request(StubKmsContext()) - self.assertEqual(len(calls), 1) - - async def test_contract_violation_surfaces_as_encryption_error(self): - # Callers see EncryptionError with ConfigurationError as its cause. - with self.assertRaises(EncryptionError) as caught: - with _wrap_encryption_errors(): - raise ConfigurationError("kms_connect_callback must return ...") - self.assertIsInstance(caught.exception.__cause__, ConfigurationError) - - async def test_network_error_from_callback_propagates(self): - async def callback(context): - raise OSError("proxy unreachable") - - # Not a ConfigurationError, so kms_request retries it. - with self.assertRaises(OSError): - await _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) - - async def test_csot_deadline_stops_a_hung_callback(self): - # A callback that ignores the timeout cannot block past the CSOT - # deadline, and a socket it yields later must be closed. - if _IS_SYNC: - raise unittest.SkipTest("the sync API cannot interrupt a callback") - - left, _right = self._socketpair() - - async def hung_callback(context): - await asyncio.sleep(0.5) - return left - - with self.assertRaises(NetworkTimeout): - with pymongo.timeout(0.1): - await _connect_kms(_KMS_ADDRESS, self._pool_options(), hung_callback, 10.0) - self.assertNotEqual(left.fileno(), -1) - # Let the shielded callback finish; the driver closes the late result. - await asyncio.sleep(0.75) - self.assertEqual(left.fileno(), -1) - - async def test_cancelling_kms_connect_closes_the_callback_socket(self): - # Cancelling during the TLS handshake must close the callback's socket, - # so a TLS proxy's relay threads wind down. - if _IS_SYNC: - raise unittest.SkipTest("cancellation is an async-only behavior") - - server_ctx = _tls_server_context() - listener = self._listen() - gate = threading.Event() - eof = threading.Event() - - def stub_server(): - conn = None - try: - conn, _ = listener.accept() - # The cancel may land before or after the executor starts the - # handshake. Peek for the ClientHello without consuming it, or - # for EOF if the driver closed it, before wrap_socket detaches conn. - while True: - data = conn.recv(4096, socket.MSG_PEEK) - if not data: - eof.set() - return - if data[:1] == b"\x16": # TLS handshake record - break - # Hold the handshake open until the test has cancelled. - if not gate.wait(5): - return - tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) - try: - tls.do_handshake() - # A discarded connection may end in a reset rather than a - # clean EOF; either proves the driver closed it. - while tls.recv(4096): - pass - except OSError: - pass - eof.set() - tls.close() - except OSError: - pass - finally: - if conn is not None: - conn.close() - - threading.Thread(target=stub_server, daemon=True).start() - - options = self._pool_options(_client_tls_context()) - socks = [] - - def connect(): - sock = socket.create_connection(listener.getsockname(), timeout=10) - socks.append(sock) - return sock - - async def callback(context): - return await _run_blocking(connect) - - # The sync version returns a socket instead of a coroutine, so both - # error codes are needed depending on the version being checked. - pending = _connect_kms(listener.getsockname(), options, callback, 10.0) - task = asyncio.ensure_future(pending) # type: ignore[type-var,arg-type] - for _ in range(100): - if socks: - break - await asyncio.sleep(0.01) - self.assertTrue(socks, "callback was never invoked") - # Bias the cancel to land mid-handshake; the stub handles the earlier - # window too. - await asyncio.sleep(0.1) - task.cancel() - with self.assertRaises(asyncio.CancelledError): - await task - # The late SSLSocket (or raw socket) must be closed; the stub sees EOF. - gate.set() - for _ in range(50): - if eof.is_set(): - break - await asyncio.sleep(0.1) - self.assertTrue(eof.is_set(), "driver never closed the callback socket") - - async def test_client_encryption_accepts_callback(self): - async def callback(context): - raise AssertionError("not called") - - client = self.simple_client() - encryption = AsyncClientEncryption( - {"local": {"key": b"\x00" * 96}}, - "keyvault.datakeys", - client, - OPTS, - kms_connect_callback=callback, - ) - self.addAsyncCleanup(encryption.close) - self.assertIs(encryption._io_callbacks.opts._kms_connect_callback, callback) - - async def test_client_encryption_rejects_non_callable(self): - client = self.simple_client() - with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): - AsyncClientEncryption( - {"local": {"key": b"\x00" * 96}}, - "keyvault.datakeys", - client, - OPTS, - kms_connect_callback="not-callable", # type: ignore[arg-type] - ) - - -class TestKmsConnectCallbackProse(AsyncEncryptionIntegrationTest): - @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") - async def asyncSetUp(self): - await super().asyncSetUp() - self.callback_calls: list[Any] = [] - - async def plain_callback(self, context): - self.callback_calls.append(context) - return await AsyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) - - def _proxy_tls_context(self): - ctx = ssl.create_default_context(cafile=CA_PEM) - ctx.check_hostname = False - # PYTHON-5040 tracks re-enabling verification once the test CA cert - # is fixed; the evergreen-tools CA lacks an Authority Key Identifier - # that newer OpenSSL requires, so verification fails on Windows 3.14. - ctx.verify_mode = ssl.CERT_NONE - return ctx - - async def tls_callback(self, context): - self.callback_calls.append(context) - callback = AsyncHTTPProxyKMSConnect( - KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, self._proxy_tls_context() - ) - return await callback(context) - - async def proxy_request(self, method, path, tls=False): - """Call the proxy's control endpoints and return the body.""" - if _IS_SYNC: - return self._proxy_request(method, path, tls) - return await asyncio.get_running_loop().run_in_executor( - None, self._proxy_request, method, path, tls - ) - - def _proxy_request(self, method, path, tls=False): - if tls: - conn = http.client.HTTPSConnection( - f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() - ) - else: - conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") - try: - conn.request(method, path) - return conn.getresponse().read().decode() - finally: - conn.close() - - async def connect_count(self, tls=False): - body = await self.proxy_request("GET", "/metrics", tls=tls) - # One "key value" per line; the server also emits connect_target. - for line in body.splitlines(): - key, _, value = line.partition(" ") - if key == "connect_count": - return int(value) - raise AssertionError(f"no connect_count in metrics body: {body!r}") - - async def test_01_plain_http_proxy(self): - await self.proxy_request("POST", "/reset") - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(await self.connect_count(), 1) - - async def test_02_https_proxy(self): - await self.proxy_request("POST", "/reset", tls=True) - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.tls_callback, - ) - await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(await self.connect_count(tls=True), 1) - - async def test_03_auto_encryption_through_proxy(self): - await self.client.keyvault.datakeys.drop() - await self.client.db.coll.drop() - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - data_key_id = await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - schema = { - "bsonType": "object", - "properties": { - "encrypted_string": { - "encrypt": { - "keyId": [data_key_id], - "bsonType": "string", - "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", - } - } - }, - } - - await self.proxy_request("POST", "/reset") - opts = AutoEncryptionOpts( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - schema_map={"db.coll": schema}, - kms_connect_callback=self.plain_callback, - ) - client_encrypted = await self.async_rs_or_single_client(auto_encryption_opts=opts) - - await client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) - decrypted = await client_encrypted.db.coll.find_one({"_id": 1}) - self.assertEqual(decrypted["encrypted_string"], "hello") - - raw = await self.client.db.coll.find_one({"_id": 1}) - self.assertIsInstance(raw["encrypted_string"], Binary) - - # The decrypt reuses the cached key, so exactly one KMS request follows - # the reset. - self.assertEqual(await self.connect_count(), 1) - - async def test_04_callback_error(self): - async def failing_callback(context): - raise OSError("proxy is on fire") - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=failing_callback, - ) - with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): - await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - - @unittest.skip( - "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " - "callback always receives the default KMS connect timeout" - ) - async def test_05_callback_receives_timeout(self): - key_vault_client = await self.async_rs_or_single_client(timeoutMS=1000) - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - key_vault_client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - - self.assertTrue(self.callback_calls, "callback was never invoked") - for context in self.callback_calls: - # Checks only the spec's non-zero requirement, which cannot fail. - self.assertIsNotNone(context.timeout) - self.assertGreater(context.timeout, 0) - - async def test_06_retry_after_network_error(self): - state = {"calls": 0} - - async def flaky_callback(context): - state["calls"] += 1 - if state["calls"] == 1: - raise OSError("first attempt fails") - return await AsyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=flaky_callback, - ) - await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(state["calls"], 2) diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index fac9a49d43..ce12200aa2 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -1,251 +1,279 @@ -"""Tests for the KMS connect callback and HTTP proxy support.""" +# Copyright 2026-present MongoDB, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for the KMS connect callback. + +This file is not processed by synchro: the tests are written once and +parameterized over both the asynchronous and synchronous APIs, selected with +the ``[async]``/``[sync]`` parametrize ids. Every test runs as a coroutine on +a pytest-asyncio loop. The ``[sync]`` variants call the blocking synchronous +APIs from within the coroutine, which is harmless for these self-contained +tests. The ``flavor`` parameter provides the per-API seams. + +The tests must run single threaded under thread based parallelization such as +pytest-run-parallel. pytest-asyncio does not support the plugin's concurrent +replicas of one test, and the tests spin up real sockets and threads that +every replica would race on. +""" from __future__ import annotations import asyncio import dataclasses -import http.client import os import socket import ssl import threading -import time -import unittest from asyncio.trsock import TransportSocket +from collections.abc import Callable +from contextlib import contextmanager from typing import Any from unittest import mock import pytest import pymongo -from bson.binary import Binary -from pymongo.encryption_options import ( - AutoEncryptionOpts, - HTTPProxyKMSConnect, - KMSConnectContext, - SyncHTTPProxyKMSConnect, -) +from bson.codec_options import CodecOptions +from pymongo.encryption_options import _HAVE_PYMONGOCRYPT, AutoEncryptionOpts, KMSConnectContext from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout from pymongo.pool_options import PoolOptions from pymongo.ssl_support import get_ssl_context -from pymongo.synchronous.encryption import ( - ClientEncryption, - _connect_kms, - _EncryptionIO, - _wrap_encryption_errors, -) -from test import PyMongoTestCase -from test.helpers_shared import AWS_CREDS, CA_PEM, CERT_PATH, CLIENT_PEM -from test.test_encryption import OPTS, EncryptionIntegrationTest +from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM -_IS_SYNC = True - -pytestmark = pytest.mark.encryption +pytestmark = [pytest.mark.encryption, pytest.mark.asyncio] _KMS_ADDRESS = ("kms.example.com", 443) -KMS_PROXY_HOST = "127.0.0.1" -KMS_PROXY_PORT = 9004 -KMS_TLS_PROXY_PORT = 9005 +OPTS = CodecOptions() -AWS_MASTER_KEY = { - "region": "us-east-1", - "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", -} +class Flavor: + """The per-API seams, shared by the tests through the ``flavor`` parameter.""" -def _tls_server_context(cert=CLIENT_PEM): - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - ctx.load_cert_chain(cert) - return ctx + def __init__(self, is_async: bool) -> None: + self.is_async = is_async + async def maybe_await(self, result: Any) -> Any: + """Await ``result`` in the asynchronous flavor (a no-op otherwise).""" + if self.is_async: + return await result + return result -def _client_tls_context(verify=False): - # verify=False matches the driver's test mode: the local certs don't verify. - if verify: - return get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) - return get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) - - -def _run_blocking(func, *args): - """Run a blocking callable off the event loop (inline in the sync version).""" - if _IS_SYNC: + async def offload(self, func: Callable[..., Any], *args: Any) -> Any: + """Run a blocking callable off the event loop (inline when synchronous).""" + if self.is_async: + return await asyncio.get_running_loop().run_in_executor(None, func, *args) return func(*args) - return asyncio.get_running_loop().run_in_executor(None, func, *args) - - -def _callback_returning(value): - """A kms_connect_callback that always produces ``value``.""" - def callback(context): - return value + def encryption(self): + """The flavor's ``encryption`` module.""" + if self.is_async: + from pymongo.asynchronous import encryption + else: + from pymongo.synchronous import encryption + return encryption - return callback + def kms_connect(self): + """The flavor's ``_kms_connect`` module.""" + if self.is_async: + from pymongo.asynchronous import _kms_connect + else: + from pymongo.synchronous import _kms_connect + return _kms_connect + async def connect(self, address, pool_options, callback, timeout): + """``_connect_kms`` for this flavor.""" + module = self.kms_connect() + if self.is_async: + return await module._connect_kms(address, pool_options, callback, timeout) + return module._connect_kms(address, pool_options, callback, timeout) -def _kms_context(host="kms.example.com", port=443, timeout=10): - """A KMSConnectContext with the defaults used throughout these tests.""" - return KMSConnectContext(host=host, port=port, timeout=timeout) + def callback(self, func): + """Adapt a non-blocking ``func(context)`` to the flavor's callback form.""" + if self.is_async: + async def callback(context): + return func(context) -def _read_http_request(conn): - """Read until the blank line that ends a CONNECT request, or None on EOF.""" - request = b"" - while b"\r\n\r\n" not in request: - chunk = conn.recv(4096) - if not chunk: - return None - request += chunk - return request + return callback + return func -def _insecure_client_context(): - # PYTHON-5040 tracks re-enabling verification: the evergreen-tools CA - # lacks an Authority Key Identifier newer OpenSSL requires. - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ctx.check_hostname = False - ctx.verify_mode = ssl.CERT_NONE - return ctx + def blocking_callback(self, func): + """Adapt a blocking ``func(context)``, offloaded in the async flavor.""" + if self.is_async: + async def callback(context): + return await asyncio.get_running_loop().run_in_executor(None, func, context) -class TestKmsConnectCallbackUnit(PyMongoTestCase): - """Contract checks for kms_connect_callback that need no KMS server.""" + return callback - @staticmethod - def _pool_options(ssl_context=None): - return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=ssl_context) + return func - def _listen(self, backlog=1): - listener = socket.socket() - listener.bind(("127.0.0.1", 0)) - listener.listen(backlog) - self.addCleanup(listener.close) - return listener + def callback_returning(self, value): + """A kms_connect_callback that always produces ``value``.""" + return self.callback(lambda context: value) - def _socketpair(self): - left, right = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) - return left, right + def client_tls_context(self, verify=False): + # verify=False matches the driver's test mode: the local certs don't verify. + if verify: + return get_ssl_context(None, None, CA_PEM, None, False, False, False, not self.is_async) + return get_ssl_context(None, None, None, None, True, True, False, not self.is_async) - def _start_proxy(self, handler, backlog=1): - """Serve each accepted connection with ``handler(conn)`` in a daemon thread.""" - listener = self._listen(backlog) + def client_encryption(self, kms_providers, key_vault_namespace, client, kms_connect_callback): + """A ClientEncryption for this flavor using a local key provider.""" + if self.is_async: + from pymongo.asynchronous.encryption import AsyncClientEncryption - def serve(): - for _ in range(backlog): - try: - conn, _ = listener.accept() - except OSError: - return - try: - handler(conn) - except OSError: - pass - finally: - conn.close() + return AsyncClientEncryption( + kms_providers, + key_vault_namespace, + client, + OPTS, + kms_connect_callback=kms_connect_callback, + ) + from pymongo.synchronous.encryption import ClientEncryption - threading.Thread(target=serve, daemon=True).start() - return listener.getsockname() + return ClientEncryption( + kms_providers, + key_vault_namespace, + client, + OPTS, + kms_connect_callback=kms_connect_callback, + ) - def _record_and_reply(self, accepted, reply): - """A proxy that records each CONNECT request, replies ``reply``, and closes.""" + def simple_client(self): + """A lazily-connecting client for this flavor.""" + if self.is_async: + from pymongo.asynchronous.mongo_client import AsyncMongoClient - def handler(conn): - request = _read_http_request(conn) - if request is None: - return - accepted.append(request) - conn.sendall(reply) + return AsyncMongoClient() + from pymongo import MongoClient - return self._start_proxy(handler) + return MongoClient() - def _tls_echo_proxy(self, delay=0): - """A TLS CONNECT proxy that replies 200, then echoes one tunneled read.""" - server_ctx = _tls_server_context() - def handler(conn): - tls = server_ctx.wrap_socket(conn, server_side=True) - request = _read_http_request(tls) - if request is None: - return - tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - # The tunneled peer speaks only after the client does, as a TLS - # server would; ``delay`` lets the reply outlast the CONNECT deadline. - if delay: - time.sleep(delay) - tls.sendall(b"echo:" + tls.recv(64)) - tls.close() +ASYNC = Flavor(is_async=True) +SYNC = Flavor(is_async=False) - return self._start_proxy(handler) +both_flavors = pytest.mark.parametrize("flavor", [ASYNC, SYNC], ids=["async", "sync"]) +async_only = pytest.mark.parametrize("flavor", [ASYNC], ids=["async"]) - def _echo_over_tunnel(self, sock): - sock.settimeout(10) - _run_blocking(sock.sendall, b"ping") - data = _run_blocking(sock.recv, 64) - self.assertEqual(data, b"echo:ping") - def test_init_kms_connect_callback(self): - opts = AutoEncryptionOpts({}, "k.d") - self.assertIsNone(opts._kms_connect_callback) +def _pool_options(ssl_context=None): + return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=ssl_context) - def callback(context): - raise AssertionError("not called") - opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) - self.assertIs(opts._kms_connect_callback, callback) +def _tls_server_context(cert=CLIENT_PEM): + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + ctx.load_cert_chain(cert) + return ctx - for bad in [1, "not-callable", object()]: - with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): - AutoEncryptionOpts({}, "k.d", kms_connect_callback=bad) # type: ignore[arg-type] - context = KMSConnectContext(host="kms.example.com", port=443, timeout=9.5) - self.assertEqual(context.host, "kms.example.com") - self.assertEqual(context.port, 443) - self.assertEqual(context.timeout, 9.5) - with self.assertRaises(dataclasses.FrozenInstanceError): - context.host = "evil.example.com" # type: ignore[misc] +@contextmanager +def _listen(backlog=1): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(backlog) + try: + yield listener + finally: + listener.close() + + +@contextmanager +def _socketpair(): + left, right = socket.socketpair() + try: + yield left, right + finally: + left.close() + right.close() + + +@both_flavors +async def test_init_kms_connect_callback(flavor): + opts = AutoEncryptionOpts({}, "k.d") + assert opts._kms_connect_callback is None + + def action(context): + raise AssertionError("not called") + + callback = flavor.callback(action) + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) + assert opts._kms_connect_callback is callback + + for bad in [1, "not-callable", object()]: + with pytest.raises(TypeError, match="kms_connect_callback must be callable"): + AutoEncryptionOpts({}, "k.d", kms_connect_callback=bad) # type: ignore[arg-type] + + context = KMSConnectContext(host="kms.example.com", port=443, timeout=9.5) + assert context.host == "kms.example.com" + assert context.port == 443 + assert context.timeout == 9.5 + with pytest.raises(dataclasses.FrozenInstanceError): + context.host = "evil.example.com" # type: ignore[misc] + + +@both_flavors +async def test_non_socket_return_raises_configuration_error(flavor): + with pytest.raises(ConfigurationError, match="must return a connected"): + await flavor.connect( + _KMS_ADDRESS, _pool_options(), flavor.callback_returning("not-a-socket"), 10.0 + ) - def test_non_socket_return_raises_configuration_error(self): - with self.assertRaisesRegex(ConfigurationError, "must return a connected"): - _connect_kms( - _KMS_ADDRESS, self._pool_options(), _callback_returning("not-a-socket"), 10.0 - ) - def test_already_wrapped_socket_is_rejected(self): - # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ctx.check_hostname = False - ctx.verify_mode = ssl.CERT_NONE - left, _right = self._socketpair() - # No peer needed to produce a genuine ssl.SSLSocket. +@both_flavors +async def test_already_wrapped_socket_is_rejected(flavor): + # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + with _socketpair() as (left, _right): + # No peer needed to produce a genuine ssl.SSLSocket. The driver closes + # the wrapped socket when rejecting it. wrapped = ctx.wrap_socket(left, do_handshake_on_connect=False, server_hostname="x") - self.addCleanup(wrapped.close) + with pytest.raises(ConfigurationError, match="unwrapped"): + await flavor.connect( + _KMS_ADDRESS, _pool_options(), flavor.callback_returning(wrapped), 10.0 + ) - with self.assertRaisesRegex(ConfigurationError, "unwrapped"): - _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(wrapped), 10.0) - def test_context_receives_host_port_and_timeout(self): - received = [] - left, _right = self._socketpair() +@both_flavors +async def test_context_receives_host_port_and_timeout(flavor): + received = [] + with _socketpair() as (left, _right): - def callback(context): + def action(context): received.append(context) return left # ssl_context=None returns the socket unchanged, so a plain socket is accepted. - conn = _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 12.5) - self.assertIs(conn, left) + conn = await flavor.connect(_KMS_ADDRESS, _pool_options(), flavor.callback(action), 12.5) + assert conn is left + + assert len(received) == 1 + assert received[0].host == "kms.example.com" + assert received[0].port == 443 + assert received[0].timeout == 12.5 - self.assertEqual(len(received), 1) - self.assertEqual(received[0].host, "kms.example.com") - self.assertEqual(received[0].port, 443) - self.assertEqual(received[0].timeout, 12.5) - def test_non_blocking_socket_from_callback_is_accepted(self): - # Without the driver normalizing the mode, this raises ValueError. - server_ctx = _tls_server_context() - listener = self._listen() +@both_flavors +async def test_non_blocking_socket_from_callback_is_accepted(flavor): + # Without the driver normalizing the mode, this raises ValueError. + server_ctx = _tls_server_context() + with _listen() as listener: def serve(): try: @@ -256,27 +284,33 @@ def serve(): threading.Thread(target=serve, daemon=True).start() - options = self._pool_options(_client_tls_context()) + options = _pool_options(flavor.client_tls_context()) def connect(): sock = socket.create_connection(listener.getsockname(), timeout=10) sock.setblocking(False) return sock - def callback(context): - return _run_blocking(connect) + conn = await flavor.connect( + listener.getsockname(), + options, + flavor.blocking_callback(lambda context: connect()), + 10.0, + ) + try: + assert conn.gettimeout() is not None + finally: + conn.close() - conn = _connect_kms(listener.getsockname(), options, callback, 10.0) - self.addCleanup(conn.close) - self.assertIsNotNone(conn.gettimeout()) - def test_tls_verification_targets_the_kms_host(self): - # The handshake must verify against the KMS address, not the peer the - # callback connected to. The server cert covers 127.0.0.1 (the peer) - # and localhost, but not the KMS hostname used below, so only - # address-based verification produces this outcome. - server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) - listener = self._listen(2) +@both_flavors +async def test_tls_verification_targets_the_kms_host(flavor): + # The handshake must verify against the KMS address, not the peer the + # callback connected to. The server cert covers 127.0.0.1 (the peer) + # and localhost, but not the KMS hostname used below, so only + # address-based verification produces this outcome. + server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) + with _listen(2) as listener: def serve(): for _ in range(2): @@ -290,7 +324,7 @@ def serve(): threading.Thread(target=serve, daemon=True).start() # Full verification: trusted CA, invalid certs and hostnames rejected. - options = self._pool_options(_client_tls_context(verify=True)) + options = _pool_options(flavor.client_tls_context(verify=True)) created = [] @@ -299,441 +333,210 @@ def connect(): created.append(sock) return sock - def callback(context): - return _run_blocking(connect) - port = listener.getsockname()[1] # The cert covers localhost: verifying against the KMS address succeeds. - conn = _connect_kms(("localhost", port), options, callback, 10.0) - self.addCleanup(conn.close) - # TLS-wrapped in either SSL flavor: a new object, not the plain socket. - self.assertIsNot(conn, created[0]) + conn = await flavor.connect( + ("localhost", port), + options, + flavor.blocking_callback(lambda context: connect()), + 10.0, + ) + try: + # TLS-wrapped in either SSL flavor: a new object, not the plain socket. + assert conn is not created[0] + finally: + conn.close() # The cert does not cover this name: verification must fail even # though the peer (127.0.0.1) presents a cert valid for itself. - with self.assertRaises(ConnectionFailure): - _connect_kms(("kms.example.com", port), options, callback, 10.0) + with pytest.raises(ConnectionFailure): + await flavor.connect( + ("kms.example.com", port), + options, + flavor.blocking_callback(lambda context: connect()), + 10.0, + ) - def test_asyncio_transport_socket_is_rejected(self): - # get_extra_info("socket") is a TransportSocket, not a socket.socket. - left, _right = self._socketpair() - with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): - _connect_kms( - _KMS_ADDRESS, self._pool_options(), _callback_returning(TransportSocket(left)), 10.0 +@both_flavors +async def test_asyncio_transport_socket_is_rejected(flavor): + # get_extra_info("socket") is a TransportSocket, not a socket.socket. + with _socketpair() as (left, _right): + with pytest.raises(ConfigurationError, match="TransportSocket"): + await flavor.connect( + _KMS_ADDRESS, + _pool_options(), + flavor.callback_returning(TransportSocket(left)), + 10.0, ) - def test_cancelled_tls_wrap_closes_late_socket(self): - # A cancelled wrap can leave the executor producing an SSLSocket; the - # done callback must close it. - if _IS_SYNC: - raise unittest.SkipTest("the cancel-safe wrap is an async path") - from pymongo.pool_shared import _close_late_socket - left, _right = self._socketpair() +@async_only +async def test_cancelled_tls_wrap_closes_late_socket(flavor): + # A cancelled wrap can leave the executor producing an SSLSocket. The + # done callback must close it. + from pymongo.pool_shared import _close_late_socket + + with _socketpair() as (left, _right): future = asyncio.get_running_loop().create_future() future.set_result(left) - self.assertNotEqual(left.fileno(), -1) + assert left.fileno() != -1 _close_late_socket(future) - self.assertEqual(left.fileno(), -1) - - def test_http_proxy_helper_tunnels_and_reports_refusal(self): - # Covers the CONNECT handshake without KMS credentials. - accepted: list[bytes] = [] - context = _kms_context() - - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - sock = SyncHTTPProxyKMSConnect(host, port)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") - - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n" - ) - with self.assertRaisesRegex(OSError, "refused CONNECT"): - SyncHTTPProxyKMSConnect(host, port)(context) - - # Any 2xx status is a successful tunnel, not just HTTP/1.1 200. - host, port = self._record_and_reply( - accepted, b"HTTP/1.0 200 Connection Established\r\n\r\n" - ) - sock = SyncHTTPProxyKMSConnect(host, port)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - - # A status code must be exactly three digits, with no zero padding. - for reply in (b"HTTP/1.1 2000 Evil\r\n\r\n", b"HTTP/1.1 00200 Evil\r\n\r\n"): - host, port = self._record_and_reply(accepted, reply) - with self.assertRaisesRegex(OSError, "refused CONNECT"): - SyncHTTPProxyKMSConnect(host, port)(context) - - def test_control_characters_in_kms_host_are_rejected(self): - # Reject CR/LF in the configurable host before it reaches CONNECT. - callback = SyncHTTPProxyKMSConnect("proxy.example.com", 8080) - context = _kms_context(host="kms.example.com\r\nX-Injected: 1") - with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): - callback(context) - # Whitespace would split the request line into extra tokens. - context = _kms_context(host="kms.example.com ") - with self.assertRaisesRegex(ConfigurationError, "control characters or whitespace"): - callback(context) - - def test_http_proxy_helper_sends_custom_headers(self): - # Extra CONNECT headers reach the proxy verbatim. - accepted: list[bytes] = [] - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Trace-Id": "abc123"} - sock = SyncHTTPProxyKMSConnect(host, port, headers=headers)(_kms_context()) - self.addCleanup(sock.close) - request = accepted[0] - self.assertEqual(request.split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") - self.assertIn(b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n", request) - self.assertIn(b"\r\nX-Trace-Id: abc123\r\n", request) - self.assertEqual(request.count(b"\r\nHost: "), 1) - - def test_http_proxy_helper_authenticates_to_the_proxy(self): - # The motivating case: 407 without credentials, 200 with them. - def handler(conn): - request = _read_http_request(conn) - if request is None: - return - if b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request: - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - else: - conn.sendall(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") - - host, port = self._start_proxy(handler, backlog=2) - context = _kms_context() - with self.assertRaisesRegex(OSError, "refused CONNECT"): - SyncHTTPProxyKMSConnect(host, port)(context) - headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz"} - sock = SyncHTTPProxyKMSConnect(host, port, headers=headers)(context) - self.addCleanup(sock.close) - self.assertIsInstance(sock, socket.socket) - - def test_http_proxy_helper_rejects_bad_headers(self): - for headers in [ - {"Bad\r\nName": "x"}, - {"Bad Name": "x"}, - {"Bad\tName": "x"}, - {"X-Ok": "ok\r\nInjected: 1"}, - {"Host": "evil.example.com"}, - {"host": "evil.example.com"}, - {"": "x"}, - {"Bad:Name": "x"}, - ]: - with self.assertRaisesRegex(ConfigurationError, "proxy header|Host CONNECT header"): - SyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) - - for headers in [{1: "x"}, {"X-Ok": 1}, {None: "x"}, {"X-Ok": None}]: - with self.assertRaisesRegex(TypeError, "must be strings"): - SyncHTTPProxyKMSConnect("proxy.example.com", 8080, headers=headers) - - def test_http_proxy_helper_accepts_legal_header_values(self): - # Colons and spaces are legal in values (e.g. auth schemes); only - # CR/LF would let a value inject a request line. - callback = SyncHTTPProxyKMSConnect( - "proxy.example.com", - 8080, - headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, - ) - self.assertEqual( - callback.headers, - {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, - ) - - def test_tls_proxy_helper_bridges_the_tunnel(self): - # Covers the TLS-proxy path and the socketpair relay without KMS creds. - host, port = self._tls_echo_proxy() - sock = HTTPProxyKMSConnect(host, port, _insecure_client_context())(_kms_context()) - self.addCleanup(sock.close) - self._echo_over_tunnel(sock) - - def test_bridge_does_not_inherit_the_connect_deadline(self): - # The relay must outlast the much shorter CONNECT deadline. - host, port = self._tls_echo_proxy(delay=3.0) - sock = HTTPProxyKMSConnect(host, port, _insecure_client_context())( - _kms_context(timeout=2.0) - ) - self.addCleanup(sock.close) - self._echo_over_tunnel(sock) - - def test_proxy_closing_before_connect_reply_raises(self): - def handler(conn): - # Read the CONNECT request, then hang up without replying. - conn.recv(4096) - - host, port = self._start_proxy(handler) - with self.assertRaisesRegex(OSError, "proxy closed the connection"): - SyncHTTPProxyKMSConnect(host, port)(_kms_context()) - - def test_cancelled_proxy_connect_closes_the_late_socket(self): - # A cancelled connect must close the socket the executor thread - # produces after the cancellation. - if _IS_SYNC: - raise unittest.SkipTest("cancellation is an async-only behavior") - - requested = threading.Event() - reply = threading.Event() - - def handler(conn): - conn.recv(4096) - requested.set() - if not reply.wait(10): - return - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - # Keep the connection open so the tunnel can complete its reads. - time.sleep(0.1) - - host, port = self._start_proxy(handler) - - tunneled: list[socket.socket] = [] - original_tunnel = HTTPProxyKMSConnect._tunnel - - def spy_tunnel(self, sock, context, deadline): - tunneled.append(sock) - original_tunnel(self, sock, context, deadline) - - with mock.patch.object(HTTPProxyKMSConnect, "_tunnel", spy_tunnel): - task = asyncio.create_task(SyncHTTPProxyKMSConnect(host, port)(_kms_context())) - waited = _run_blocking(requested.wait, 10) - self.assertTrue(waited, "proxy never received the CONNECT request") - task.cancel("no longer needed") - with self.assertRaises(asyncio.CancelledError): - task - # Let the stub reply, completing the executor's future late. - reply.set() - time.sleep(0.5) - - self.assertEqual(len(tunneled), 1) - self.assertEqual(tunneled[0].fileno(), -1, "late socket was left open") - - def test_connect_timeout_is_not_reclassified(self): - # A connect that times out keeps its socket.timeout type instead of - # being reported as a generic connect error. - def timeout_connect(self, address): - raise socket.timeout("timed out") - - with mock.patch.object(socket.socket, "connect", timeout_connect): - with self.assertRaises(socket.timeout): - HTTPProxyKMSConnect("127.0.0.1", 9999)._connect_proxy(time.monotonic() + 10) - - def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): - # A proxy may coalesce its 200 with tunneled bytes; reading past the header would drop them. - def handler(conn): - conn.recv(4096) - conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") - - host, port = self._start_proxy(handler) - sock = SyncHTTPProxyKMSConnect(host, port)(_kms_context()) - self.addCleanup(sock.close) - sock.settimeout(10) - data = _run_blocking(sock.recv, 64) - self.assertEqual(data, b"early-bytes") - - def test_ipv6_host_is_bracketed_in_connect(self): - accepted: list[bytes] = [] - host, port = self._record_and_reply( - accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n" - ) - sock = SyncHTTPProxyKMSConnect(host, port)(_kms_context(host="::1")) - self.addCleanup(sock.close) - self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") - - def test_oversized_connect_response_is_rejected(self): - def handler(conn): - conn.recv(4096) - # Never sends the terminator. - while True: - conn.sendall(b"x" * 1024) - - host, port = self._start_proxy(handler) - with self.assertRaisesRegex(OSError, "oversized CONNECT response"): - SyncHTTPProxyKMSConnect(host, port)(_kms_context()) + assert left.fileno() == -1 - def test_remaining_raises_once_the_deadline_passes(self): - from pymongo.encryption_options import _remaining - self.assertGreater(_remaining(time.monotonic() + 5), 0) - with self.assertRaises(socket.timeout): - _remaining(time.monotonic() - 1) +@async_only +async def test_non_coroutine_callback_is_rejected(flavor): + # A plain def must be rejected before it blocks the event loop. + entered = [] - def test_bridge_failure_closes_the_proxy_socket(self): - # A failure inside _bridge must not strand the connected proxy socket. - server_ctx = _tls_server_context() - - def handler(conn): - tls = server_ctx.wrap_socket(conn, server_side=True) - tls.recv(4096) - tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") - tls.close() - - captured = [] - - def failing_bridge(self, proxy): - captured.append(proxy) - raise OSError("no file descriptors") - - host, port = self._start_proxy(handler) - context = _kms_context() - - with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): - with self.assertRaisesRegex(OSError, "no file descriptors"): - HTTPProxyKMSConnect(host, port, _insecure_client_context())(context) + def callback(context): + entered.append(context) + return None - self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") + with pytest.raises(ConfigurationError, match="coroutine function"): + await flavor.connect(_KMS_ADDRESS, _pool_options(), callback, 10.0) + assert entered == [], "invalid callback must not be entered" - def test_non_coroutine_callback_is_rejected(self): - # A plain def must be rejected before it blocks the event loop. - if _IS_SYNC: - raise unittest.SkipTest("a regular function is correct for the sync API") - entered = [] +@both_flavors +async def test_unconnected_socket_from_callback_is_rejected(flavor): + # An unconnected socket would fail later as a transient error and be retried. + with socket.socket() as bare: + with pytest.raises(ConfigurationError, match="already connected"): + await flavor.connect( + _KMS_ADDRESS, _pool_options(), flavor.callback_returning(bare), 10.0 + ) - def callback(context): - entered.append(context) - return None - with self.assertRaisesRegex(ConfigurationError, "coroutine function"): - _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) - self.assertEqual(entered, [], "invalid callback must not be entered") +@both_flavors +async def test_datagram_socket_from_callback_is_rejected(flavor): + # TLS on a connected UDP socket raises NotImplementedError, which would be retried. + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as left: + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as right: + right.bind(("127.0.0.1", 0)) + left.connect(right.getsockname()) - def test_unconnected_socket_from_callback_is_rejected(self): - # An unconnected socket would fail later as a transient error and be retried. - bare = socket.socket() - self.addCleanup(bare.close) + with pytest.raises(ConfigurationError, match="stream socket"): + await flavor.connect( + _KMS_ADDRESS, _pool_options(), flavor.callback_returning(left), 10.0 + ) - with self.assertRaisesRegex(ConfigurationError, "already connected"): - _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(bare), 10.0) - def test_datagram_socket_from_callback_is_rejected(self): - # TLS on a connected UDP socket raises NotImplementedError, which would be retried. - left = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - right = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - self.addCleanup(left.close) - self.addCleanup(right.close) - right.bind(("127.0.0.1", 0)) - left.connect(right.getsockname()) +@both_flavors +async def test_kms_request_does_not_retry_a_contract_violation(flavor): + # _connect_kms has no retry loop. The no-retry guarantee is in + # kms_request, so exercise that instead. + calls = [] - with self.assertRaisesRegex(ConfigurationError, "stream socket"): - _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(left), 10.0) + def action(context): + calls.append(context) + return "not-a-socket" - def test_kms_request_does_not_retry_a_contract_violation(self): - # _connect_kms has no retry loop; the no-retry guarantee is in - # kms_request, so exercise that instead. - calls = [] + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=flavor.callback(action)) + io = flavor.encryption()._EncryptionIO(None, mock.MagicMock(), None, opts) - def callback(context): - calls.append(context) - return "not-a-socket" + class StubKmsContext: + endpoint = "kms.example.com:443" + message = b"request" + kms_provider = "aws" + usleep = 0 + bytes_needed = 1 - opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) - io = _EncryptionIO(None, mock.MagicMock(), None, opts) + def feed(self, data): + raise AssertionError("should not reach the socket") - class StubKmsContext: - endpoint = "kms.example.com:443" - message = b"request" - kms_provider = "aws" - usleep = 0 - bytes_needed = 1 + def fail(self): + raise AssertionError("a contract violation must not be retried") - def feed(self, data): - raise AssertionError("should not reach the socket") + with pytest.raises(ConfigurationError): + await flavor.maybe_await(io.kms_request(StubKmsContext())) + assert len(calls) == 1 - def fail(self): - raise AssertionError("a contract violation must not be retried") - with self.assertRaises(ConfigurationError): - io.kms_request(StubKmsContext()) - self.assertEqual(len(calls), 1) +@both_flavors +async def test_contract_violation_surfaces_as_encryption_error(flavor): + # Callers see EncryptionError with ConfigurationError as its cause. + with pytest.raises(EncryptionError) as exc_info: + with flavor.encryption()._wrap_encryption_errors(): + raise ConfigurationError("kms_connect_callback must return ...") + assert isinstance(exc_info.value.__cause__, ConfigurationError) - def test_contract_violation_surfaces_as_encryption_error(self): - # Callers see EncryptionError with ConfigurationError as its cause. - with self.assertRaises(EncryptionError) as caught: - with _wrap_encryption_errors(): - raise ConfigurationError("kms_connect_callback must return ...") - self.assertIsInstance(caught.exception.__cause__, ConfigurationError) - def test_network_error_from_callback_propagates(self): - def callback(context): - raise OSError("proxy unreachable") +@both_flavors +async def test_network_error_from_callback_propagates(flavor): + def action(context): + raise OSError("proxy unreachable") - # Not a ConfigurationError, so kms_request retries it. - with self.assertRaises(OSError): - _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) + # Not a ConfigurationError, so kms_request retries it. + with pytest.raises(OSError): + await flavor.connect(_KMS_ADDRESS, _pool_options(), flavor.callback(action), 10.0) - def test_csot_deadline_stops_a_hung_callback(self): - # A callback that ignores the timeout cannot block past the CSOT - # deadline, and a socket it yields later must be closed. - if _IS_SYNC: - raise unittest.SkipTest("the sync API cannot interrupt a callback") - left, _right = self._socketpair() +@async_only +async def test_csot_deadline_stops_a_hung_callback(flavor): + # A callback that ignores the timeout cannot block past the CSOT + # deadline, and a socket it yields later must be closed. + with _socketpair() as (left, _right): - def hung_callback(context): - time.sleep(0.5) + async def hung_callback(context): + await asyncio.sleep(0.5) return left - with self.assertRaises(NetworkTimeout): + with pytest.raises(NetworkTimeout): with pymongo.timeout(0.1): - _connect_kms(_KMS_ADDRESS, self._pool_options(), hung_callback, 10.0) - self.assertNotEqual(left.fileno(), -1) - # Let the shielded callback finish; the driver closes the late result. - time.sleep(0.75) - self.assertEqual(left.fileno(), -1) - - def test_cancelling_kms_connect_closes_the_callback_socket(self): - # Cancelling during the TLS handshake must close the callback's socket, - # so a TLS proxy's relay threads wind down. - if _IS_SYNC: - raise unittest.SkipTest("cancellation is an async-only behavior") - - server_ctx = _tls_server_context() - listener = self._listen() - gate = threading.Event() - eof = threading.Event() - - def stub_server(): - conn = None - try: - conn, _ = listener.accept() - # The cancel may land before or after the executor starts the - # handshake. Peek for the ClientHello without consuming it, or - # for EOF if the driver closed it, before wrap_socket detaches conn. - while True: - data = conn.recv(4096, socket.MSG_PEEK) - if not data: - eof.set() - return - if data[:1] == b"\x16": # TLS handshake record - break - # Hold the handshake open until the test has cancelled. - if not gate.wait(5): + await flavor.connect(_KMS_ADDRESS, _pool_options(), hung_callback, 10.0) + assert left.fileno() != -1 + # Let the shielded callback finish. The driver closes the late result. + await asyncio.sleep(0.75) + assert left.fileno() == -1 + + +@async_only +async def test_cancelling_kms_connect_closes_the_callback_socket(flavor): + # Cancelling during the TLS handshake must close the callback's socket, + # so a TLS proxy's relay threads wind down. + server_ctx = _tls_server_context() + gate = threading.Event() + eof = threading.Event() + + def stub_server(listener): + conn = None + try: + conn, _ = listener.accept() + # The cancel may land before or after the executor starts the + # handshake. Peek for the ClientHello without consuming it, or + # for EOF if the driver closed it, before wrap_socket detaches conn. + while True: + data = conn.recv(4096, socket.MSG_PEEK) + if not data: + eof.set() return - tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) - try: - tls.do_handshake() - # A discarded connection may end in a reset rather than a - # clean EOF; either proves the driver closed it. - while tls.recv(4096): - pass - except OSError: + if data[:1] == b"\x16": # TLS handshake record + break + # Hold the handshake open until the test has cancelled. + if not gate.wait(5): + return + tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) + try: + tls.do_handshake() + # A discarded connection may end in a reset rather than a + # clean EOF. Either outcome proves the driver closed it. + while tls.recv(4096): pass - eof.set() - tls.close() except OSError: pass - finally: - if conn is not None: - conn.close() + eof.set() + tls.close() + except OSError: + pass + finally: + if conn is not None: + conn.close() - threading.Thread(target=stub_server, daemon=True).start() + with _listen() as listener: + threading.Thread(target=stub_server, args=(listener,), daemon=True).start() - options = self._pool_options(_client_tls_context()) + options = _pool_options(flavor.client_tls_context()) socks = [] def connect(): @@ -741,234 +544,63 @@ def connect(): socks.append(sock) return sock - def callback(context): - return _run_blocking(connect) - - # The sync version returns a socket instead of a coroutine, so both - # error codes are needed depending on the version being checked. - pending = _connect_kms(listener.getsockname(), options, callback, 10.0) - task = asyncio.ensure_future(pending) # type: ignore[type-var,arg-type] + # Schedule the connect now and cancel it once the callback socket + # exists, so the cancellation lands mid-handshake. + task = asyncio.ensure_future( + flavor.connect( + listener.getsockname(), + options, + flavor.blocking_callback(lambda context: connect()), + 10.0, + ) + ) for _ in range(100): if socks: break - time.sleep(0.01) - self.assertTrue(socks, "callback was never invoked") - # Bias the cancel to land mid-handshake; the stub handles the earlier + await asyncio.sleep(0.01) + assert socks, "callback was never invoked" + # Bias the cancel to land mid-handshake. The stub handles the earlier # window too. - time.sleep(0.1) + await asyncio.sleep(0.1) task.cancel() - with self.assertRaises(asyncio.CancelledError): - task - # The late SSLSocket (or raw socket) must be closed; the stub sees EOF. + with pytest.raises(asyncio.CancelledError): + await task + # The late SSLSocket (or raw socket) must be closed. The stub sees EOF. gate.set() for _ in range(50): if eof.is_set(): break - time.sleep(0.1) - self.assertTrue(eof.is_set(), "driver never closed the callback socket") - - def test_client_encryption_accepts_callback(self): - def callback(context): - raise AssertionError("not called") - - client = self.simple_client() - encryption = ClientEncryption( - {"local": {"key": b"\x00" * 96}}, - "keyvault.datakeys", - client, - OPTS, - kms_connect_callback=callback, - ) - self.addCleanup(encryption.close) - self.assertIs(encryption._io_callbacks.opts._kms_connect_callback, callback) - - def test_client_encryption_rejects_non_callable(self): - client = self.simple_client() - with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): - ClientEncryption( - {"local": {"key": b"\x00" * 96}}, - "keyvault.datakeys", - client, - OPTS, - kms_connect_callback="not-callable", # type: ignore[arg-type] - ) + await asyncio.sleep(0.1) + assert eof.is_set(), "driver never closed the callback socket" -class TestKmsConnectCallbackProse(EncryptionIntegrationTest): - @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") - def setUp(self): - super().setUp() - self.callback_calls: list[Any] = [] - - def plain_callback(self, context): - self.callback_calls.append(context) - return SyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) - - def _proxy_tls_context(self): - ctx = ssl.create_default_context(cafile=CA_PEM) - ctx.check_hostname = False - # PYTHON-5040 tracks re-enabling verification once the test CA cert - # is fixed; the evergreen-tools CA lacks an Authority Key Identifier - # that newer OpenSSL requires, so verification fails on Windows 3.14. - ctx.verify_mode = ssl.CERT_NONE - return ctx - - def tls_callback(self, context): - self.callback_calls.append(context) - callback = SyncHTTPProxyKMSConnect( - KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, self._proxy_tls_context() - ) - return callback(context) - - def proxy_request(self, method, path, tls=False): - """Call the proxy's control endpoints and return the body.""" - if _IS_SYNC: - return self._proxy_request(method, path, tls) - return asyncio.get_running_loop().run_in_executor( - None, self._proxy_request, method, path, tls - ) +@pytest.mark.skipif(not _HAVE_PYMONGOCRYPT, reason="pymongocrypt is not installed") +@both_flavors +async def test_client_encryption_accepts_callback(flavor): + def action(context): + raise AssertionError("not called") - def _proxy_request(self, method, path, tls=False): - if tls: - conn = http.client.HTTPSConnection( - f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() - ) - else: - conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") - try: - conn.request(method, path) - return conn.getresponse().read().decode() - finally: - conn.close() - - def connect_count(self, tls=False): - body = self.proxy_request("GET", "/metrics", tls=tls) - # One "key value" per line; the server also emits connect_target. - for line in body.splitlines(): - key, _, value = line.partition(" ") - if key == "connect_count": - return int(value) - raise AssertionError(f"no connect_count in metrics body: {body!r}") - - def test_01_plain_http_proxy(self): - self.proxy_request("POST", "/reset") - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(self.connect_count(), 1) - - def test_02_https_proxy(self): - self.proxy_request("POST", "/reset", tls=True) - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.tls_callback, - ) - encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(self.connect_count(tls=True), 1) - - def test_03_auto_encryption_through_proxy(self): - self.client.keyvault.datakeys.drop() - self.client.db.coll.drop() - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - data_key_id = encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - schema = { - "bsonType": "object", - "properties": { - "encrypted_string": { - "encrypt": { - "keyId": [data_key_id], - "bsonType": "string", - "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", - } - } - }, - } - - self.proxy_request("POST", "/reset") - opts = AutoEncryptionOpts( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - schema_map={"db.coll": schema}, - kms_connect_callback=self.plain_callback, - ) - client_encrypted = self.rs_or_single_client(auto_encryption_opts=opts) - - client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) - decrypted = client_encrypted.db.coll.find_one({"_id": 1}) - self.assertEqual(decrypted["encrypted_string"], "hello") - - raw = self.client.db.coll.find_one({"_id": 1}) - self.assertIsInstance(raw["encrypted_string"], Binary) - - # The decrypt reuses the cached key, so exactly one KMS request follows - # the reset. - self.assertEqual(self.connect_count(), 1) - - def test_04_callback_error(self): - def failing_callback(context): - raise OSError("proxy is on fire") - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=failing_callback, - ) - with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): - encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - - @unittest.skip( - "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " - "callback always receives the default KMS connect timeout" + callback = flavor.callback(action) + client = flavor.simple_client() + encryption = flavor.client_encryption( + {"local": {"key": b"\x00" * 96}}, "keyvault.datakeys", client, callback ) - def test_05_callback_receives_timeout(self): - key_vault_client = self.rs_or_single_client(timeoutMS=1000) - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, - "keyvault.datakeys", - key_vault_client, - OPTS, - kms_connect_callback=self.plain_callback, - ) - encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - - self.assertTrue(self.callback_calls, "callback was never invoked") - for context in self.callback_calls: - # Checks only the spec's non-zero requirement, which cannot fail. - self.assertIsNotNone(context.timeout) - self.assertGreater(context.timeout, 0) - - def test_06_retry_after_network_error(self): - state = {"calls": 0} - - def flaky_callback(context): - state["calls"] += 1 - if state["calls"] == 1: - raise OSError("first attempt fails") - return SyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) - - encryption = self.create_client_encryption( - {"aws": AWS_CREDS}, + try: + assert encryption._io_callbacks.opts._kms_connect_callback is callback + finally: + await flavor.maybe_await(encryption.close()) + await flavor.maybe_await(client.close()) + + +@pytest.mark.skipif(not _HAVE_PYMONGOCRYPT, reason="pymongocrypt is not installed") +@both_flavors +async def test_client_encryption_rejects_non_callable(flavor): + client = flavor.simple_client() + with pytest.raises(TypeError, match="kms_connect_callback must be callable"): + flavor.client_encryption( + {"local": {"key": b"\x00" * 96}}, "keyvault.datakeys", - self.client, - OPTS, - kms_connect_callback=flaky_callback, + client, + "not-callable", ) - encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) - self.assertGreaterEqual(state["calls"], 2) + await flavor.maybe_await(client.close()) From 748af090386290a20bfeee984a40f46db52203b6 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Thu, 8 Oct 2026 17:18:13 -0500 Subject: [PATCH 3/4] PYTHON-6154 Fix the KMS connect refactor Deliver the PR's stated goals: - Rename pymongo/_kms_connect.py to pymongo/_kms_connect_shared.py, so the shared module reads naturally next to the per-API pymongo/{asynchronous,synchronous}/_kms_connect.py modules. - Genericize the shared module docstring: it enumerated the module's exact contents, which would go stale as helpers move in. - Drop the PYTHON-6147 feature tests (HTTP proxy helpers and the prose class) that leaked into this refactor and broke test collection: they import HTTPProxyKMSConnect/AsyncHTTPProxyKMSConnect, which are added by PYTHON-6147, not by this PR. - Replace the synchro-mirrored test pair with a single hand-written test/test_kms_connect.py, written once and parameterized over both APIs through a Facade class (the 18 pre-existing tests, 32 variants, byte-identical bodies). The async-side file is gone, so synchro no longer mirrors it. - Tighten the docstrings and comments across the touched modules and the test file. --- pymongo/_kms_connect_shared.py | 12 +- pymongo/asynchronous/_kms_connect.py | 20 +-- pymongo/synchronous/_kms_connect.py | 20 +-- test/test_kms_connect.py | 214 +++++++++++++-------------- 4 files changed, 128 insertions(+), 138 deletions(-) diff --git a/pymongo/_kms_connect_shared.py b/pymongo/_kms_connect_shared.py index 710f380d92..c160657da9 100644 --- a/pymongo/_kms_connect_shared.py +++ b/pymongo/_kms_connect_shared.py @@ -38,8 +38,7 @@ class KMSConnectContext: :param host: Hostname of the KMS server, and the TLS verification target. :param port: Port of the KMS server. :param timeout: Seconds allowed for the connection: the default KMS - connect timeout, capped by the remaining time of an active operation - timeout (``timeoutMS``). + connect timeout, capped by the remaining ``timeoutMS`` budget. .. note:: ``timeoutMS`` does not constrain KMS requests for explicit encryption, so ``timeout`` is always the default there. Automatic @@ -62,9 +61,8 @@ class KMSConnectContext: def _close_rejected_kms_socket(obj: Any) -> None: """Close a callback return value that failed validation, best effort. - ``_connect_kms`` raises on a contract violation instead of returning the - value, so no caller ever takes ownership of it. Close it here, tolerating - non-socket values and ``close()`` failures. + ``_connect_kms`` raises on a contract violation, so no caller takes + ownership. Tolerates non-socket values and ``close()`` failures. """ close = getattr(obj, "close", None) if callable(close): @@ -72,6 +70,6 @@ def _close_rejected_kms_socket(obj: Any) -> None: close() -# Sphinx documents this class under pymongo.encryption_options, the public -# import path, so the definition must claim that module name. +# Sphinx documents this class under pymongo.encryption_options, so the +# definition must claim that module name. KMSConnectContext.__module__ = "pymongo.encryption_options" diff --git a/pymongo/asynchronous/_kms_connect.py b/pymongo/asynchronous/_kms_connect.py index 3f4ba1a7d3..7e256286a9 100644 --- a/pymongo/asynchronous/_kms_connect.py +++ b/pymongo/asynchronous/_kms_connect.py @@ -72,9 +72,9 @@ async def _connect_kms( except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - # TLS targets address, not the peer, so verification follows the KMS host. - # A plain callable would block the event loop before we could reject it, - # so check the callback first. + # TLS targets ``address``, not the peer, so verification follows the KMS + # host. Reject plain callables up front: invoking one would block the + # event loop. if not _IS_SYNC: callback_any: Any = kms_connect_callback is_coro = inspect.iscoroutinefunction(callback_any) @@ -90,13 +90,13 @@ async def _connect_kms( ) remaining = _csot.remaining() if remaining is None or _IS_SYNC: - # The synchronous API cannot interrupt a callback that has started - # running; honoring the deadline is the callback's contract there. + # The synchronous API cannot interrupt a running callback; honoring + # the deadline is the callback's contract there. sock = await result else: - # CSOT is cooperative: a callback that ignores the timeout could block - # past the deadline. Shield the task so stopping the wait does not - # cancel it mid-flight, and close any socket it yields later. + # CSOT is cooperative: a callback that ignores the timeout can + # outlive the deadline. Shield the task so cancelling the wait does + # not cancel it mid-flight, and close the socket it yields later. task = asyncio.ensure_future(result) try: sock = await asyncio.wait_for(asyncio.shield(task), remaining) @@ -130,8 +130,8 @@ async def _connect_kms( "kms_connect_callback must return a stream socket, not a datagram one." ) # The callback may have consumed much of the CSOT budget, and wrapping - # resets the socket timeout, so recompute the remaining time here and for - # the KMS request that follows. + # resets the socket timeout, so recompute the remaining time for the KMS + # request that follows. sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) try: conn = await _async_wrap_socket_tls(sock, address, opts) diff --git a/pymongo/synchronous/_kms_connect.py b/pymongo/synchronous/_kms_connect.py index b05ba019a4..aca683a2d5 100644 --- a/pymongo/synchronous/_kms_connect.py +++ b/pymongo/synchronous/_kms_connect.py @@ -72,9 +72,9 @@ def _connect_kms( except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) - # TLS targets address, not the peer, so verification follows the KMS host. - # A plain callable would block the event loop before we could reject it, - # so check the callback first. + # TLS targets ``address``, not the peer, so verification follows the KMS + # host. Reject plain callables up front: invoking one would block the + # event loop. if not _IS_SYNC: callback_any: Any = kms_connect_callback is_coro = inspect.iscoroutinefunction(callback_any) @@ -90,13 +90,13 @@ def _connect_kms( ) remaining = _csot.remaining() if remaining is None or _IS_SYNC: - # The synchronous API cannot interrupt a callback that has started - # running; honoring the deadline is the callback's contract there. + # The synchronous API cannot interrupt a running callback; honoring + # the deadline is the callback's contract there. sock = result else: - # CSOT is cooperative: a callback that ignores the timeout could block - # past the deadline. Shield the task so stopping the wait does not - # cancel it mid-flight, and close any socket it yields later. + # CSOT is cooperative: a callback that ignores the timeout can + # outlive the deadline. Shield the task so cancelling the wait does + # not cancel it mid-flight, and close the socket it yields later. task = asyncio.ensure_future(result) try: sock = asyncio.wait_for(asyncio.shield(task), remaining) @@ -130,8 +130,8 @@ def _connect_kms( "kms_connect_callback must return a stream socket, not a datagram one." ) # The callback may have consumed much of the CSOT budget, and wrapping - # resets the socket timeout, so recompute the remaining time here and for - # the KMS request that follows. + # resets the socket timeout, so recompute the remaining time for the KMS + # request that follows. sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) try: conn = _wrap_socket_tls(sock, address, opts) diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index ce12200aa2..8209651776 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -14,17 +14,16 @@ """Unit tests for the KMS connect callback. -This file is not processed by synchro: the tests are written once and -parameterized over both the asynchronous and synchronous APIs, selected with -the ``[async]``/``[sync]`` parametrize ids. Every test runs as a coroutine on -a pytest-asyncio loop. The ``[sync]`` variants call the blocking synchronous -APIs from within the coroutine, which is harmless for these self-contained -tests. The ``flavor`` parameter provides the per-API seams. - -The tests must run single threaded under thread based parallelization such as -pytest-run-parallel. pytest-asyncio does not support the plugin's concurrent -replicas of one test, and the tests spin up real sockets and threads that -every replica would race on. +Not processed by synchro: the tests are written once and parameterized over +the asynchronous and synchronous APIs with the ``[async]``/``[sync]`` ids. +Every test runs as a coroutine on a pytest-asyncio loop; the ``[sync]`` +variants call the blocking synchronous APIs inline, which is harmless for +these self-contained tests. The ``api`` parameter provides the per-API +accessors. + +The tests must run single threaded under thread based parallelization such +as pytest-run-parallel: pytest-asyncio does not support concurrent replicas, +and the tests spin up real sockets and threads that replicas would race on. """ from __future__ import annotations @@ -58,14 +57,14 @@ OPTS = CodecOptions() -class Flavor: - """The per-API seams, shared by the tests through the ``flavor`` parameter.""" +class Facade: + """The per-API accessors, shared by the tests through the ``api`` parameter.""" def __init__(self, is_async: bool) -> None: self.is_async = is_async async def maybe_await(self, result: Any) -> Any: - """Await ``result`` in the asynchronous flavor (a no-op otherwise).""" + """Await ``result`` in the asynchronous API (a no-op otherwise).""" if self.is_async: return await result return result @@ -77,7 +76,7 @@ async def offload(self, func: Callable[..., Any], *args: Any) -> Any: return func(*args) def encryption(self): - """The flavor's ``encryption`` module.""" + """The API's ``encryption`` module.""" if self.is_async: from pymongo.asynchronous import encryption else: @@ -85,7 +84,7 @@ def encryption(self): return encryption def kms_connect(self): - """The flavor's ``_kms_connect`` module.""" + """The API's ``_kms_connect`` module.""" if self.is_async: from pymongo.asynchronous import _kms_connect else: @@ -93,14 +92,14 @@ def kms_connect(self): return _kms_connect async def connect(self, address, pool_options, callback, timeout): - """``_connect_kms`` for this flavor.""" + """``_connect_kms`` for this API.""" module = self.kms_connect() if self.is_async: return await module._connect_kms(address, pool_options, callback, timeout) return module._connect_kms(address, pool_options, callback, timeout) def callback(self, func): - """Adapt a non-blocking ``func(context)`` to the flavor's callback form.""" + """Adapt a non-blocking ``func(context)`` to the API's callback form.""" if self.is_async: async def callback(context): @@ -111,7 +110,7 @@ async def callback(context): return func def blocking_callback(self, func): - """Adapt a blocking ``func(context)``, offloaded in the async flavor.""" + """Adapt a blocking ``func(context)``, offloaded in the async API.""" if self.is_async: async def callback(context): @@ -132,7 +131,7 @@ def client_tls_context(self, verify=False): return get_ssl_context(None, None, None, None, True, True, False, not self.is_async) def client_encryption(self, kms_providers, key_vault_namespace, client, kms_connect_callback): - """A ClientEncryption for this flavor using a local key provider.""" + """A ClientEncryption for this API using a local key provider.""" if self.is_async: from pymongo.asynchronous.encryption import AsyncClientEncryption @@ -154,7 +153,7 @@ def client_encryption(self, kms_providers, key_vault_namespace, client, kms_conn ) def simple_client(self): - """A lazily-connecting client for this flavor.""" + """A lazily-connecting client for this API.""" if self.is_async: from pymongo.asynchronous.mongo_client import AsyncMongoClient @@ -164,11 +163,11 @@ def simple_client(self): return MongoClient() -ASYNC = Flavor(is_async=True) -SYNC = Flavor(is_async=False) +ASYNC = Facade(is_async=True) +SYNC = Facade(is_async=False) -both_flavors = pytest.mark.parametrize("flavor", [ASYNC, SYNC], ids=["async", "sync"]) -async_only = pytest.mark.parametrize("flavor", [ASYNC], ids=["async"]) +both_apis = pytest.mark.parametrize("api", [ASYNC, SYNC], ids=["async", "sync"]) +async_only = pytest.mark.parametrize("api", [ASYNC], ids=["async"]) def _pool_options(ssl_context=None): @@ -202,15 +201,15 @@ def _socketpair(): right.close() -@both_flavors -async def test_init_kms_connect_callback(flavor): +@both_apis +async def test_init_kms_connect_callback(api): opts = AutoEncryptionOpts({}, "k.d") assert opts._kms_connect_callback is None def action(context): raise AssertionError("not called") - callback = flavor.callback(action) + callback = api.callback(action) opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) assert opts._kms_connect_callback is callback @@ -226,32 +225,29 @@ def action(context): context.host = "evil.example.com" # type: ignore[misc] -@both_flavors -async def test_non_socket_return_raises_configuration_error(flavor): +@both_apis +async def test_non_socket_return_raises_configuration_error(api): with pytest.raises(ConfigurationError, match="must return a connected"): - await flavor.connect( - _KMS_ADDRESS, _pool_options(), flavor.callback_returning("not-a-socket"), 10.0 + await api.connect( + _KMS_ADDRESS, _pool_options(), api.callback_returning("not-a-socket"), 10.0 ) -@both_flavors -async def test_already_wrapped_socket_is_rejected(flavor): +@both_apis +async def test_already_wrapped_socket_is_rejected(api): # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) ctx.check_hostname = False ctx.verify_mode = ssl.CERT_NONE with _socketpair() as (left, _right): - # No peer needed to produce a genuine ssl.SSLSocket. The driver closes - # the wrapped socket when rejecting it. + # No peer is needed to produce a genuine ssl.SSLSocket. wrapped = ctx.wrap_socket(left, do_handshake_on_connect=False, server_hostname="x") with pytest.raises(ConfigurationError, match="unwrapped"): - await flavor.connect( - _KMS_ADDRESS, _pool_options(), flavor.callback_returning(wrapped), 10.0 - ) + await api.connect(_KMS_ADDRESS, _pool_options(), api.callback_returning(wrapped), 10.0) -@both_flavors -async def test_context_receives_host_port_and_timeout(flavor): +@both_apis +async def test_context_receives_host_port_and_timeout(api): received = [] with _socketpair() as (left, _right): @@ -260,7 +256,7 @@ def action(context): return left # ssl_context=None returns the socket unchanged, so a plain socket is accepted. - conn = await flavor.connect(_KMS_ADDRESS, _pool_options(), flavor.callback(action), 12.5) + conn = await api.connect(_KMS_ADDRESS, _pool_options(), api.callback(action), 12.5) assert conn is left assert len(received) == 1 @@ -269,8 +265,8 @@ def action(context): assert received[0].timeout == 12.5 -@both_flavors -async def test_non_blocking_socket_from_callback_is_accepted(flavor): +@both_apis +async def test_non_blocking_socket_from_callback_is_accepted(api): # Without the driver normalizing the mode, this raises ValueError. server_ctx = _tls_server_context() with _listen() as listener: @@ -284,17 +280,17 @@ def serve(): threading.Thread(target=serve, daemon=True).start() - options = _pool_options(flavor.client_tls_context()) + options = _pool_options(api.client_tls_context()) def connect(): sock = socket.create_connection(listener.getsockname(), timeout=10) sock.setblocking(False) return sock - conn = await flavor.connect( + conn = await api.connect( listener.getsockname(), options, - flavor.blocking_callback(lambda context: connect()), + api.blocking_callback(lambda context: connect()), 10.0, ) try: @@ -303,12 +299,11 @@ def connect(): conn.close() -@both_flavors -async def test_tls_verification_targets_the_kms_host(flavor): +@both_apis +async def test_tls_verification_targets_the_kms_host(api): # The handshake must verify against the KMS address, not the peer the - # callback connected to. The server cert covers 127.0.0.1 (the peer) - # and localhost, but not the KMS hostname used below, so only - # address-based verification produces this outcome. + # callback connected to: the cert covers 127.0.0.1 and localhost, but + # not the KMS hostname used below. server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) with _listen(2) as listener: @@ -324,7 +319,7 @@ def serve(): threading.Thread(target=serve, daemon=True).start() # Full verification: trusted CA, invalid certs and hostnames rejected. - options = _pool_options(flavor.client_tls_context(verify=True)) + options = _pool_options(api.client_tls_context(verify=True)) created = [] @@ -335,43 +330,43 @@ def connect(): port = listener.getsockname()[1] # The cert covers localhost: verifying against the KMS address succeeds. - conn = await flavor.connect( + conn = await api.connect( ("localhost", port), options, - flavor.blocking_callback(lambda context: connect()), + api.blocking_callback(lambda context: connect()), 10.0, ) try: - # TLS-wrapped in either SSL flavor: a new object, not the plain socket. + # TLS-wrapped in either API: a new object, not the plain socket. assert conn is not created[0] finally: conn.close() - # The cert does not cover this name: verification must fail even - # though the peer (127.0.0.1) presents a cert valid for itself. + # The cert does not cover this name, so verification fails even + # though the peer's cert is valid for itself. with pytest.raises(ConnectionFailure): - await flavor.connect( + await api.connect( ("kms.example.com", port), options, - flavor.blocking_callback(lambda context: connect()), + api.blocking_callback(lambda context: connect()), 10.0, ) -@both_flavors -async def test_asyncio_transport_socket_is_rejected(flavor): +@both_apis +async def test_asyncio_transport_socket_is_rejected(api): # get_extra_info("socket") is a TransportSocket, not a socket.socket. with _socketpair() as (left, _right): with pytest.raises(ConfigurationError, match="TransportSocket"): - await flavor.connect( + await api.connect( _KMS_ADDRESS, _pool_options(), - flavor.callback_returning(TransportSocket(left)), + api.callback_returning(TransportSocket(left)), 10.0, ) @async_only -async def test_cancelled_tls_wrap_closes_late_socket(flavor): +async def test_cancelled_tls_wrap_closes_late_socket(api): # A cancelled wrap can leave the executor producing an SSLSocket. The # done callback must close it. from pymongo.pool_shared import _close_late_socket @@ -385,7 +380,7 @@ async def test_cancelled_tls_wrap_closes_late_socket(flavor): @async_only -async def test_non_coroutine_callback_is_rejected(flavor): +async def test_non_coroutine_callback_is_rejected(api): # A plain def must be rejected before it blocks the event loop. entered = [] @@ -394,22 +389,20 @@ def callback(context): return None with pytest.raises(ConfigurationError, match="coroutine function"): - await flavor.connect(_KMS_ADDRESS, _pool_options(), callback, 10.0) + await api.connect(_KMS_ADDRESS, _pool_options(), callback, 10.0) assert entered == [], "invalid callback must not be entered" -@both_flavors -async def test_unconnected_socket_from_callback_is_rejected(flavor): +@both_apis +async def test_unconnected_socket_from_callback_is_rejected(api): # An unconnected socket would fail later as a transient error and be retried. with socket.socket() as bare: with pytest.raises(ConfigurationError, match="already connected"): - await flavor.connect( - _KMS_ADDRESS, _pool_options(), flavor.callback_returning(bare), 10.0 - ) + await api.connect(_KMS_ADDRESS, _pool_options(), api.callback_returning(bare), 10.0) -@both_flavors -async def test_datagram_socket_from_callback_is_rejected(flavor): +@both_apis +async def test_datagram_socket_from_callback_is_rejected(api): # TLS on a connected UDP socket raises NotImplementedError, which would be retried. with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as left: with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as right: @@ -417,14 +410,12 @@ async def test_datagram_socket_from_callback_is_rejected(flavor): left.connect(right.getsockname()) with pytest.raises(ConfigurationError, match="stream socket"): - await flavor.connect( - _KMS_ADDRESS, _pool_options(), flavor.callback_returning(left), 10.0 - ) + await api.connect(_KMS_ADDRESS, _pool_options(), api.callback_returning(left), 10.0) -@both_flavors -async def test_kms_request_does_not_retry_a_contract_violation(flavor): - # _connect_kms has no retry loop. The no-retry guarantee is in +@both_apis +async def test_kms_request_does_not_retry_a_contract_violation(api): + # _connect_kms has no retry loop; the no-retry guarantee is in # kms_request, so exercise that instead. calls = [] @@ -432,8 +423,8 @@ def action(context): calls.append(context) return "not-a-socket" - opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=flavor.callback(action)) - io = flavor.encryption()._EncryptionIO(None, mock.MagicMock(), None, opts) + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=api.callback(action)) + io = api.encryption()._EncryptionIO(None, mock.MagicMock(), None, opts) class StubKmsContext: endpoint = "kms.example.com:443" @@ -449,31 +440,31 @@ def fail(self): raise AssertionError("a contract violation must not be retried") with pytest.raises(ConfigurationError): - await flavor.maybe_await(io.kms_request(StubKmsContext())) + await api.maybe_await(io.kms_request(StubKmsContext())) assert len(calls) == 1 -@both_flavors -async def test_contract_violation_surfaces_as_encryption_error(flavor): +@both_apis +async def test_contract_violation_surfaces_as_encryption_error(api): # Callers see EncryptionError with ConfigurationError as its cause. with pytest.raises(EncryptionError) as exc_info: - with flavor.encryption()._wrap_encryption_errors(): + with api.encryption()._wrap_encryption_errors(): raise ConfigurationError("kms_connect_callback must return ...") assert isinstance(exc_info.value.__cause__, ConfigurationError) -@both_flavors -async def test_network_error_from_callback_propagates(flavor): +@both_apis +async def test_network_error_from_callback_propagates(api): def action(context): raise OSError("proxy unreachable") # Not a ConfigurationError, so kms_request retries it. with pytest.raises(OSError): - await flavor.connect(_KMS_ADDRESS, _pool_options(), flavor.callback(action), 10.0) + await api.connect(_KMS_ADDRESS, _pool_options(), api.callback(action), 10.0) @async_only -async def test_csot_deadline_stops_a_hung_callback(flavor): +async def test_csot_deadline_stops_a_hung_callback(api): # A callback that ignores the timeout cannot block past the CSOT # deadline, and a socket it yields later must be closed. with _socketpair() as (left, _right): @@ -484,7 +475,7 @@ async def hung_callback(context): with pytest.raises(NetworkTimeout): with pymongo.timeout(0.1): - await flavor.connect(_KMS_ADDRESS, _pool_options(), hung_callback, 10.0) + await api.connect(_KMS_ADDRESS, _pool_options(), hung_callback, 10.0) assert left.fileno() != -1 # Let the shielded callback finish. The driver closes the late result. await asyncio.sleep(0.75) @@ -492,7 +483,7 @@ async def hung_callback(context): @async_only -async def test_cancelling_kms_connect_closes_the_callback_socket(flavor): +async def test_cancelling_kms_connect_closes_the_callback_socket(api): # Cancelling during the TLS handshake must close the callback's socket, # so a TLS proxy's relay threads wind down. server_ctx = _tls_server_context() @@ -504,8 +495,9 @@ def stub_server(listener): try: conn, _ = listener.accept() # The cancel may land before or after the executor starts the - # handshake. Peek for the ClientHello without consuming it, or - # for EOF if the driver closed it, before wrap_socket detaches conn. + # handshake: peek for a ClientHello without consuming it, or for + # EOF if the driver closed the socket, before wrap_socket + # detaches conn. while True: data = conn.recv(4096, socket.MSG_PEEK) if not data: @@ -519,8 +511,8 @@ def stub_server(listener): tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) try: tls.do_handshake() - # A discarded connection may end in a reset rather than a - # clean EOF. Either outcome proves the driver closed it. + # A discarded connection may reset instead of reaching clean + # EOF; either proves the driver closed it. while tls.recv(4096): pass except OSError: @@ -536,7 +528,7 @@ def stub_server(listener): with _listen() as listener: threading.Thread(target=stub_server, args=(listener,), daemon=True).start() - options = _pool_options(flavor.client_tls_context()) + options = _pool_options(api.client_tls_context()) socks = [] def connect(): @@ -547,10 +539,10 @@ def connect(): # Schedule the connect now and cancel it once the callback socket # exists, so the cancellation lands mid-handshake. task = asyncio.ensure_future( - flavor.connect( + api.connect( listener.getsockname(), options, - flavor.blocking_callback(lambda context: connect()), + api.blocking_callback(lambda context: connect()), 10.0, ) ) @@ -575,32 +567,32 @@ def connect(): @pytest.mark.skipif(not _HAVE_PYMONGOCRYPT, reason="pymongocrypt is not installed") -@both_flavors -async def test_client_encryption_accepts_callback(flavor): +@both_apis +async def test_client_encryption_accepts_callback(api): def action(context): raise AssertionError("not called") - callback = flavor.callback(action) - client = flavor.simple_client() - encryption = flavor.client_encryption( + callback = api.callback(action) + client = api.simple_client() + encryption = api.client_encryption( {"local": {"key": b"\x00" * 96}}, "keyvault.datakeys", client, callback ) try: assert encryption._io_callbacks.opts._kms_connect_callback is callback finally: - await flavor.maybe_await(encryption.close()) - await flavor.maybe_await(client.close()) + await api.maybe_await(encryption.close()) + await api.maybe_await(client.close()) @pytest.mark.skipif(not _HAVE_PYMONGOCRYPT, reason="pymongocrypt is not installed") -@both_flavors -async def test_client_encryption_rejects_non_callable(flavor): - client = flavor.simple_client() +@both_apis +async def test_client_encryption_rejects_non_callable(api): + client = api.simple_client() with pytest.raises(TypeError, match="kms_connect_callback must be callable"): - flavor.client_encryption( + api.client_encryption( {"local": {"key": b"\x00" * 96}}, "keyvault.datakeys", client, "not-callable", ) - await flavor.maybe_await(client.close()) + await api.maybe_await(client.close()) From 7a9232e3d1ce1e1ac1013cfc47e2588c1f960511 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Thu, 8 Oct 2026 17:20:09 -0500 Subject: [PATCH 4/4] PYTHON-6154 Fix the KMS connect refactor Deliver the PR's stated goals: - Rename pymongo/_kms_connect.py to pymongo/_kms_connect_shared.py, so the shared module reads naturally next to the per-API pymongo/{asynchronous,synchronous}/_kms_connect.py modules. - Genericize the shared module docstring: it enumerated the module's exact contents, which would go stale as helpers move in. - Drop the PYTHON-6147 feature tests (HTTP proxy helpers and the prose class) that leaked into this refactor and broke test collection: they import HTTPProxyKMSConnect/AsyncHTTPProxyKMSConnect, which are added by PYTHON-6147, not by this PR. - Replace the synchro-mirrored test pair with a single hand-written test/test_kms_connect.py, written once and parameterized over both APIs through a Facade class (the 18 pre-existing tests, 32 variants, byte-identical bodies). The async-side file is gone, so synchro no longer mirrors it. - Tighten the docstrings and comments across the touched modules and the test file. --- pymongo/asynchronous/_kms_connect.py | 5 ++--- pymongo/synchronous/_kms_connect.py | 5 ++--- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/pymongo/asynchronous/_kms_connect.py b/pymongo/asynchronous/_kms_connect.py index 7e256286a9..168d63425d 100644 --- a/pymongo/asynchronous/_kms_connect.py +++ b/pymongo/asynchronous/_kms_connect.py @@ -15,9 +15,8 @@ """KMS connection for the asynchronous API. Connects to a KMS host, directly or through a ``kms_connect_callback``, and -performs the KMS TLS handshake over the connection. The synchronous mirror of -this module is generated by synchro. The helpers shared by both APIs live in -``pymongo._kms_connect_shared``. +performs the KMS TLS handshake over the connection. The helpers shared by +both APIs live in ``pymongo._kms_connect_shared``. """ from __future__ import annotations diff --git a/pymongo/synchronous/_kms_connect.py b/pymongo/synchronous/_kms_connect.py index aca683a2d5..12a8da93dc 100644 --- a/pymongo/synchronous/_kms_connect.py +++ b/pymongo/synchronous/_kms_connect.py @@ -15,9 +15,8 @@ """KMS connection for the synchronous API. Connects to a KMS host, directly or through a ``kms_connect_callback``, and -performs the KMS TLS handshake over the connection. The synchronous mirror of -this module is generated by synchro. The helpers shared by both APIs live in -``pymongo._kms_connect_shared``. +performs the KMS TLS handshake over the connection. The helpers shared by +both APIs live in ``pymongo._kms_connect_shared``. """ from __future__ import annotations