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
Original file line number Diff line number Diff line change
Expand Up @@ -247,7 +247,7 @@ async def get_all(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> AsyncGenerator[DocumentSnapshot, Any]:
) -> AsyncGenerator[DocumentSnapshot[AsyncDocumentReference], Any]:
"""Retrieve a batch of documents.

.. note::
Expand Down Expand Up @@ -385,13 +385,13 @@ async def _recursive_delete(
num_deleted: int = 0

if isinstance(reference, AsyncCollectionReference):
chunk: List[DocumentSnapshot]
chunk: List[DocumentSnapshot[AsyncDocumentReference]]
async for chunk in (
reference.recursive()
.select([FieldPath.document_id()])
._chunkify(chunk_size)
):
doc_snap: DocumentSnapshot
doc_snap: DocumentSnapshot[AsyncDocumentReference]
for doc_snap in chunk:
num_deleted += 1
bulk_writer.delete(doc_snap.reference)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -207,7 +207,7 @@ async def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[AsyncDocumentReference]]:
"""Read the documents in this collection.

This sends a ``RunQuery`` RPC and returns a list of documents
Expand Down Expand Up @@ -254,7 +254,7 @@ def stream(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> AsyncStreamGenerator[DocumentSnapshot]:
) -> AsyncStreamGenerator[DocumentSnapshot[AsyncDocumentReference]]:
"""Read the documents in this collection.

This sends a ``RunQuery`` RPC and then returns a generator which
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -333,7 +333,7 @@ async def get(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> DocumentSnapshot:
) -> DocumentSnapshot[AsyncDocumentReference]:
"""Retrieve a snapshot of the current document.

See :meth:`~google.cloud.firestore_v1.base_client.BaseClient.field_path` for
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
import google.cloud.firestore_v1.types.query_profile as query_profile_pb

# Types needed only for Type Hints
from google.cloud.firestore_v1.async_document import AsyncDocumentReference
from google.cloud.firestore_v1.async_transaction import AsyncTransaction
from google.cloud.firestore_v1.base_document import DocumentSnapshot
from google.cloud.firestore_v1.base_vector_query import DistanceMeasure
Expand Down Expand Up @@ -152,11 +153,11 @@ def __init__(

async def _chunkify(
self, chunk_size: int
) -> AsyncGenerator[List[DocumentSnapshot], None]:
) -> AsyncGenerator[List[DocumentSnapshot[AsyncDocumentReference]], None]:
max_to_return: Optional[int] = self._limit
num_returned: int = 0
original: AsyncQuery = self._copy()
last_document: Optional[DocumentSnapshot] = None
last_document: Optional[DocumentSnapshot[AsyncDocumentReference]] = None

while True:
# Optionally trim the `chunk_size` down to honor a previously
Expand Down Expand Up @@ -196,7 +197,7 @@ async def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[AsyncDocumentReference]]:
"""Read the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and returns a list of documents
Expand Down Expand Up @@ -356,7 +357,9 @@ async def _make_stream(
timeout: Optional[float] = None,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> AsyncGenerator[DocumentSnapshot | query_profile_pb.ExplainMetrics, Any]:
) -> AsyncGenerator[
DocumentSnapshot[AsyncDocumentReference] | query_profile_pb.ExplainMetrics, Any
]:
"""Internal method for stream(). Read the documents in the collection
that match this query.

Expand Down Expand Up @@ -438,7 +441,7 @@ def stream(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> AsyncStreamGenerator[DocumentSnapshot]:
) -> AsyncStreamGenerator[DocumentSnapshot[AsyncDocumentReference]]:
"""Read the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and then returns a generator which
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -174,7 +174,7 @@ async def get_all(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> AsyncGenerator[DocumentSnapshot, Any]:
) -> AsyncGenerator[DocumentSnapshot[AsyncDocumentReference], Any]:
"""Retrieves multiple documents from Firestore.

Args:
Expand Down Expand Up @@ -206,7 +206,10 @@ async def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> AsyncGenerator[DocumentSnapshot, Any] | AsyncStreamGenerator[DocumentSnapshot]:
) -> (
AsyncGenerator[DocumentSnapshot[AsyncDocumentReference], Any]
| AsyncStreamGenerator[DocumentSnapshot[AsyncDocumentReference]]
):
"""
Retrieve a document or a query result from the database.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
if TYPE_CHECKING: # pragma: NO COVER
import google.cloud.firestore_v1.types.query_profile as query_profile_pb
from google.cloud.firestore_v1 import transaction
from google.cloud.firestore_v1.async_document import AsyncDocumentReference
from google.cloud.firestore_v1.base_document import DocumentSnapshot
from google.cloud.firestore_v1.query_profile import ExplainMetrics, ExplainOptions

Expand All @@ -58,7 +59,7 @@ async def get(
timeout: Optional[float] = None,
*,
explain_options: Optional[ExplainOptions] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[AsyncDocumentReference]]:
"""Runs the vector query.

This sends a ``RunQuery`` RPC and returns a list of document messages.
Expand Down Expand Up @@ -109,7 +110,9 @@ async def _make_stream(
retry: retries.AsyncRetry | object | None = gapic_v1.method.DEFAULT,
timeout: Optional[float] = None,
explain_options: Optional[ExplainOptions] = None,
) -> AsyncGenerator[DocumentSnapshot | query_profile_pb.ExplainMetrics, Any]:
) -> AsyncGenerator[
DocumentSnapshot[AsyncDocumentReference] | query_profile_pb.ExplainMetrics, Any
]:
"""Internal method for stream(). Read the documents in the collection
that match this query.

Expand Down Expand Up @@ -177,7 +180,7 @@ def stream(
timeout: Optional[float] = None,
*,
explain_options: Optional[ExplainOptions] = None,
) -> AsyncStreamGenerator[DocumentSnapshot]:
) -> AsyncStreamGenerator[DocumentSnapshot[AsyncDocumentReference]]:
"""Reads the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and then returns an iterator which
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,11 @@
Any,
Awaitable,
Dict,
Generic,
Iterable,
Optional,
Tuple,
TypeVar,
Union,
)

Expand Down Expand Up @@ -359,7 +361,10 @@ def on_snapshot(self, callback):
raise NotImplementedError


class DocumentSnapshot(object):
DocRefType = TypeVar("DocRefType", bound=BaseDocumentReference)


class DocumentSnapshot(Generic[DocRefType]):
Comment thread
daniel-sanche marked this conversation as resolved.
"""A snapshot of document data in a Firestore database.

This represents data retrieved at a specific time and may not contain
Expand All @@ -371,7 +376,7 @@ class DocumentSnapshot(object):
:meth:`~google.cloud.DocumentReference.get`.

Args:
reference (:class:`~google.cloud.firestore_v1.document.DocumentReference`):
reference (Union[:class:`~google.cloud.firestore_v1.document.DocumentReference`, :class:`~google.cloud.firestore_v1.async_document.AsyncDocumentReference`]):
A document reference corresponding to the document that contains
the data in this snapshot.
data (Dict[str, Any]):
Expand All @@ -388,9 +393,15 @@ class DocumentSnapshot(object):
"""

def __init__(
self, reference, data, exists, read_time, create_time, update_time
self,
reference: DocRefType,
data,
exists,
read_time,
create_time,
update_time,
) -> None:
self._reference = reference
self._reference: DocRefType = reference
# We want immutable data, so callers can't modify this value
# out from under us.
self._data = copy.deepcopy(data)
Expand Down Expand Up @@ -439,11 +450,11 @@ def id(self) -> str:
return self._reference.id

@property
def reference(self) -> BaseDocumentReference:
def reference(self) -> DocRefType:
"""Document reference corresponding to document that owns this data.

Returns:
:class:`~google.cloud.firestore_v1.document.DocumentReference`:
Union[:class:`~google.cloud.firestore_v1.document.DocumentReference`, :class:`~google.cloud.firestore_v1.async_document.AsyncDocumentReference`]:
A document reference corresponding to this document.
"""
return self._reference
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,7 @@ def get_all(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> Generator[DocumentSnapshot, Any, None]:
) -> Generator[DocumentSnapshot[DocumentReference], Any, None]:
"""Retrieve a batch of documents.

.. note::
Expand Down Expand Up @@ -353,13 +353,13 @@ def _recursive_delete(
num_deleted: int = 0

if isinstance(reference, CollectionReference):
chunk: List[DocumentSnapshot]
chunk: List[DocumentSnapshot[DocumentReference]]
for chunk in (
reference.recursive()
.select([FieldPath.document_id()])
._chunkify(chunk_size)
):
doc_snap: DocumentSnapshot
doc_snap: DocumentSnapshot[DocumentReference]
for doc_snap in chunk:
num_deleted += 1
bulk_writer.delete(doc_snap.reference)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -202,7 +202,7 @@ def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[DocumentReference]]:
"""Read the documents in this collection.

This sends a ``RunQuery`` RPC and returns a list of documents
Expand Down Expand Up @@ -249,7 +249,7 @@ def stream(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> StreamGenerator[DocumentSnapshot]:
) -> StreamGenerator[DocumentSnapshot[DocumentReference]]:
"""Read the documents in this collection.

This sends a ``RunQuery`` RPC and then returns an iterator which
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -369,7 +369,7 @@ def get(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> DocumentSnapshot:
) -> DocumentSnapshot[DocumentReference]:
"""Retrieve a snapshot of the current document.

See :meth:`~google.cloud.firestore_v1.base_client.BaseClient.field_path` for
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@
import datetime

from google.cloud.firestore_v1.base_vector_query import DistanceMeasure
from google.cloud.firestore_v1.document import DocumentReference
from google.cloud.firestore_v1.field_path import FieldPath
from google.cloud.firestore_v1.query_profile import ExplainMetrics, ExplainOptions

Expand Down Expand Up @@ -155,7 +156,7 @@ def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[DocumentReference]]:
"""Read the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and returns a list of documents
Expand Down Expand Up @@ -221,11 +222,11 @@ def get(

def _chunkify(
self, chunk_size: int
) -> Generator[List[DocumentSnapshot], None, None]:
) -> Generator[List[DocumentSnapshot[DocumentReference]], None, None]:
max_to_return: Optional[int] = self._limit
num_returned: int = 0
original: Query = self._copy()
last_document: Optional[DocumentSnapshot] = None
last_document: Optional[DocumentSnapshot[DocumentReference]] = None

while True:
# Optionally trim the `chunk_size` down to honor a previously
Expand Down Expand Up @@ -372,7 +373,7 @@ def _make_stream(
timeout: float | None = None,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> Generator[DocumentSnapshot, Any, Optional[ExplainMetrics]]:
) -> Generator[DocumentSnapshot[DocumentReference], Any, Optional[ExplainMetrics]]:
"""Internal method for stream(). Read the documents in the collection
that match this query.

Expand Down Expand Up @@ -474,7 +475,7 @@ def stream(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> StreamGenerator[DocumentSnapshot]:
) -> StreamGenerator[DocumentSnapshot[DocumentReference]]:
"""Read the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and then returns a generator which
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ def get_all(
timeout: float | None = None,
*,
read_time: datetime.datetime | None = None,
) -> Generator[DocumentSnapshot, Any, None]:
) -> Generator[DocumentSnapshot[DocumentReference], Any, None]:
"""Retrieves multiple documents from Firestore.

Args:
Expand Down Expand Up @@ -191,7 +191,10 @@ def get(
*,
explain_options: Optional[ExplainOptions] = None,
read_time: Optional[datetime.datetime] = None,
) -> StreamGenerator[DocumentSnapshot] | Generator[DocumentSnapshot, Any, None]:
) -> (
StreamGenerator[DocumentSnapshot[DocumentReference]]
| Generator[DocumentSnapshot[DocumentReference], Any, None]
):
"""Retrieve a document or a query result from the database.

Args:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
if TYPE_CHECKING: # pragma: NO COVER
from google.cloud.firestore_v1 import transaction
from google.cloud.firestore_v1.base_document import DocumentSnapshot
from google.cloud.firestore_v1.document import DocumentReference
from google.cloud.firestore_v1.query_profile import ExplainMetrics, ExplainOptions


Expand All @@ -60,7 +61,7 @@ def get(
timeout: Optional[float] = None,
*,
explain_options: Optional[ExplainOptions] = None,
) -> QueryResultsList[DocumentSnapshot]:
) -> QueryResultsList[DocumentSnapshot[DocumentReference]]:
"""Runs the vector query.

This sends a ``RunQuery`` RPC and returns a list of document messages.
Expand Down Expand Up @@ -124,7 +125,7 @@ def _make_stream(
retry: retries.Retry | object | None = gapic_v1.method.DEFAULT,
timeout: Optional[float] = None,
explain_options: Optional[ExplainOptions] = None,
) -> Generator[DocumentSnapshot, Any, Optional[ExplainMetrics]]:
) -> Generator[DocumentSnapshot[DocumentReference], Any, Optional[ExplainMetrics]]:
"""Reads the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and then returns a generator which
Expand Down Expand Up @@ -195,7 +196,7 @@ def stream(
timeout: Optional[float] = None,
*,
explain_options: Optional[ExplainOptions] = None,
) -> StreamGenerator[DocumentSnapshot]:
) -> StreamGenerator[DocumentSnapshot[DocumentReference]]:
"""Reads the documents in the collection that match this query.

This sends a ``RunQuery`` RPC and then returns a generator which
Expand Down
Loading
Loading