Skip to content
Merged
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
1 change: 1 addition & 0 deletions doc/api/pymongo/asynchronous/change_stream.rst
Original file line number Diff line number Diff line change
Expand Up @@ -4,3 +4,4 @@

.. automodule:: pymongo.asynchronous.change_stream
:members:
:inherited-members:
1 change: 1 addition & 0 deletions doc/api/pymongo/change_stream.rst
Original file line number Diff line number Diff line change
Expand Up @@ -3,3 +3,4 @@

.. automodule:: pymongo.change_stream
:members:
:inherited-members:
214 changes: 27 additions & 187 deletions pymongo/asynchronous/change_stream.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,50 +16,39 @@

from __future__ import annotations

import copy
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Generic, Optional, Union

from bson import CodecOptions, _bson_to_dict
from bson.raw_bson import RawBSONDocument
from bson.timestamp import Timestamp
from pymongo import _csot, common
from typing import TYPE_CHECKING, Any, Optional

from bson import _bson_to_dict
from pymongo import _csot
from pymongo.asynchronous.aggregation import (
_AggregationCommand,
_CollectionAggregationCommand,
_DatabaseAggregationCommand,
)
from pymongo.asynchronous.command_cursor import AsyncCommandCursor
from pymongo.collation import validate_collation_or_none
from pymongo.change_stream_shared import (
_AgnosticChangeStream,
_AgnosticClusterChangeStream,
_AgnosticCollectionChangeStream,
_AgnosticDatabaseChangeStream,
_resumable,
)
from pymongo.errors import (
ConnectionFailure,
CursorNotFound,
InvalidOperation,
OperationFailure,
PyMongoError,
)
from pymongo.operations import _Op
from pymongo.typings import _CollationIn, _DocumentType, _Pipeline
from pymongo.typings import _DocumentType

_IS_SYNC = False

if TYPE_CHECKING:
from pymongo.asynchronous.client_session import AsyncClientSession
from pymongo.asynchronous.collection import AsyncCollection
from pymongo.asynchronous.database import AsyncDatabase
from pymongo.asynchronous.mongo_client import AsyncMongoClient


def _resumable(exc: PyMongoError) -> bool:
"""Return True if given a resumable change stream error."""
if isinstance(exc, (ConnectionFailure, CursorNotFound)):
return True
if isinstance(exc, OperationFailure):
return exc.has_error_label("ResumableChangeStreamError")
return False
from pymongo.asynchronous.mongo_client import AsyncMongoClient # noqa: F401


class AsyncChangeStream(Generic[_DocumentType]):
class AsyncChangeStream(_AgnosticChangeStream[_DocumentType, "AsyncMongoClient[_DocumentType]"]):
"""The internal abstract base class for change stream cursors.

Should not be called directly by application developers. Use
Expand All @@ -71,137 +60,12 @@ class AsyncChangeStream(Generic[_DocumentType]):
.. seealso:: The MongoDB documentation on `changeStreams <https://mongodb.com/docs/manual/changeStreams/>`_.
"""

def __init__(
self,
target: Union[
AsyncMongoClient[_DocumentType],
AsyncDatabase[_DocumentType],
AsyncCollection[_DocumentType],
],
pipeline: Optional[_Pipeline],
full_document: Optional[str],
resume_after: Optional[Mapping[str, Any]],
max_await_time_ms: Optional[int],
batch_size: Optional[int],
collation: Optional[_CollationIn],
start_at_operation_time: Optional[Timestamp],
session: Optional[AsyncClientSession],
start_after: Optional[Mapping[str, Any]],
comment: Optional[Any] = None,
full_document_before_change: Optional[str] = None,
show_expanded_events: Optional[bool] = None,
) -> None:
if pipeline is None:
pipeline = []
pipeline = common.validate_list("pipeline", pipeline)
common.validate_string_or_none("full_document", full_document)
validate_collation_or_none(collation)
common.validate_non_negative_integer_or_none("batchSize", batch_size)

self._decode_custom = False
self._orig_codec_options: CodecOptions[_DocumentType] = target.codec_options
if target.codec_options.type_registry._decoder_map:
self._decode_custom = True
# Keep the type registry so that we support encoding custom types
# in the pipeline.
self._target = target.with_options( # type: ignore
codec_options=target.codec_options.with_options(document_class=RawBSONDocument)
)
else:
self._target = target

self._pipeline = copy.deepcopy(pipeline)
self._full_document = full_document
self._full_document_before_change = full_document_before_change
self._uses_start_after = start_after is not None
self._uses_resume_after = resume_after is not None
self._resume_token = copy.deepcopy(start_after or resume_after)
self._max_await_time_ms = max_await_time_ms
self._batch_size = batch_size
self._collation = collation
self._start_at_operation_time = start_at_operation_time
self._session = session
self._comment = comment
self._closed = False
self._timeout = self._target._timeout
self._show_expanded_events = show_expanded_events
_session: Optional[AsyncClientSession]

async def _initialize_cursor(self) -> None:
# Initialize cursor.
self._cursor = await self._create_cursor()

@property
def _aggregation_command_class(self) -> type[_AggregationCommand]:
"""The aggregation command class to be used."""
raise NotImplementedError

@property
def _client(self) -> AsyncMongoClient: # type: ignore[type-arg]
"""The client against which the aggregation commands for
this AsyncChangeStream will be run.
"""
raise NotImplementedError

def _change_stream_options(self) -> dict[str, Any]:
"""Return the options dict for the $changeStream pipeline stage."""
options: dict[str, Any] = {}
if self._full_document is not None:
options["fullDocument"] = self._full_document

if self._full_document_before_change is not None:
options["fullDocumentBeforeChange"] = self._full_document_before_change

resume_token = self.resume_token
if resume_token is not None:
if self._uses_start_after:
options["startAfter"] = resume_token
else:
options["resumeAfter"] = resume_token

elif self._start_at_operation_time is not None:
options["startAtOperationTime"] = self._start_at_operation_time

if self._show_expanded_events:
options["showExpandedEvents"] = self._show_expanded_events

return options

def _command_options(self) -> dict[str, Any]:
"""Return the options dict for the aggregation command."""
options = {}
if self._max_await_time_ms is not None:
options["maxAwaitTimeMS"] = self._max_await_time_ms
if self._batch_size is not None:
options["batchSize"] = self._batch_size
return options

def _aggregation_pipeline(self) -> list[dict[str, Any]]:
"""Return the full aggregation pipeline for this AsyncChangeStream."""
options = self._change_stream_options()
full_pipeline: list[dict[str, Any]] = [{"$changeStream": options}]
full_pipeline.extend(self._pipeline)
return full_pipeline

def _process_result(self, result: Mapping[str, Any]) -> None:
"""Callback that caches the postBatchResumeToken or
startAtOperationTime from a changeStream aggregate command response
containing an empty batch of change documents.
"""
if not result["cursor"]["firstBatch"]:
if "postBatchResumeToken" in result["cursor"]:
self._resume_token = result["cursor"]["postBatchResumeToken"]
elif (
self._start_at_operation_time is None
and self._uses_resume_after is False
and self._uses_start_after is False
):
self._start_at_operation_time = result.get("operationTime")
# PYTHON-2181: informative error on missing operationTime.
if self._start_at_operation_time is None:
raise OperationFailure(
f"Expected field 'operationTime' missing from command response : {result!r}"
)

async def _run_aggregation_cmd(
self, session: Optional[AsyncClientSession]
) -> AsyncCommandCursor: # type: ignore[type-arg]
Expand Down Expand Up @@ -243,15 +107,6 @@ async def close(self) -> None:
def __aiter__(self) -> AsyncChangeStream[_DocumentType]:
return self

@property
def resume_token(self) -> Optional[Mapping[str, Any]]:
"""The cached resume token that will be used to resume after the most
recently returned change.

.. versionadded:: 3.9
"""
return copy.deepcopy(self._resume_token)

@_csot.apply
async def next(self) -> _DocumentType:
"""Advance the cursor.
Expand Down Expand Up @@ -294,17 +149,6 @@ async def next(self) -> _DocumentType:

__anext__ = next

@property
def alive(self) -> bool:
"""Does this cursor have the potential to return more data?

.. note:: Even if :attr:`alive` is ``True``, :meth:`next` can raise
:exc:`StopIteration` and :meth:`try_next` can return ``None``.

.. versionadded:: 3.8
"""
return not self._closed

@_csot.apply
async def try_next(self) -> Optional[_DocumentType]:
"""Advance the cursor without blocking indefinitely.
Expand Down Expand Up @@ -408,7 +252,10 @@ async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None:
await self.close()


class AsyncCollectionChangeStream(AsyncChangeStream[_DocumentType]):
class AsyncCollectionChangeStream(
AsyncChangeStream[_DocumentType],
_AgnosticCollectionChangeStream[_DocumentType, "AsyncMongoClient[_DocumentType]"],
):
"""A change stream that watches changes on a single collection.

Should not be called directly by application developers. Use
Expand All @@ -423,12 +270,11 @@ class AsyncCollectionChangeStream(AsyncChangeStream[_DocumentType]):
def _aggregation_command_class(self) -> type[_CollectionAggregationCommand]:
return _CollectionAggregationCommand

@property
def _client(self) -> AsyncMongoClient[_DocumentType]:
return self._target.database.client


class AsyncDatabaseChangeStream(AsyncChangeStream[_DocumentType]):
class AsyncDatabaseChangeStream(
AsyncChangeStream[_DocumentType],
_AgnosticDatabaseChangeStream[_DocumentType, "AsyncMongoClient[_DocumentType]"],
):
"""A change stream that watches changes on all collections in a database.

Should not be called directly by application developers. Use
Expand All @@ -443,21 +289,15 @@ class AsyncDatabaseChangeStream(AsyncChangeStream[_DocumentType]):
def _aggregation_command_class(self) -> type[_DatabaseAggregationCommand]:
return _DatabaseAggregationCommand

@property
def _client(self) -> AsyncMongoClient[_DocumentType]:
return self._target.client


class AsyncClusterChangeStream(AsyncDatabaseChangeStream[_DocumentType]):
class AsyncClusterChangeStream(
AsyncDatabaseChangeStream[_DocumentType],
_AgnosticClusterChangeStream[_DocumentType, "AsyncMongoClient[_DocumentType]"],
):
"""A change stream that watches changes on all collections in the cluster.

Should not be called directly by application developers. Use
helper method :meth:`pymongo.asynchronous.mongo_client.AsyncMongoClient.watch` instead.

.. versionadded:: 3.7
"""

def _change_stream_options(self) -> dict[str, Any]:
options = super()._change_stream_options()
options["allChangesForCluster"] = True
return options
Loading
Loading