Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions pymongo/_kms_connect_shared.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
# 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 support shared by the synchronous and asynchronous APIs.
The per-API connection logic 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 ``timeoutMS`` budget.
.. 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, so no caller takes
ownership. Tolerates 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, so the
# definition must claim that module name.
KMSConnectContext.__module__ = "pymongo.encryption_options"
145 changes: 145 additions & 0 deletions pymongo/asynchronous/_kms_connect.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,145 @@
# 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 helpers shared by
both APIs live in ``pymongo._kms_connect_shared``.
"""

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_shared 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. 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)
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 running callback; honoring
# the deadline is the callback's contract there.
sock = await result
else:
# 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)
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 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
119 changes: 2 additions & 117 deletions pymongo/asynchronous/encryption.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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_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
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
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand All @@ -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,
Expand Down
Loading
Loading