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 @@ -78,6 +78,7 @@ replacements:
BSONObjectId,
BSONRegex,
BSONTimestamp,
BSONType,
)
from google.cloud.firestore_v1.client import Client
from google.cloud.firestore_v1.collection import CollectionReference
Expand Down Expand Up @@ -188,6 +189,7 @@ replacements:
"BSONObjectId",
"BSONRegex",
"BSONTimestamp",
"BSONType",
"Client",
"CountAggregation",
"CollectionGroup",
Expand Down Expand Up @@ -268,6 +270,7 @@ replacements:
BSONObjectId,
BSONRegex,
BSONTimestamp,
BSONType,
Client,
CollectionGroup,
CollectionReference,
Expand Down Expand Up @@ -333,6 +336,7 @@ replacements:
"BSONObjectId",
"BSONRegex",
"BSONTimestamp",
"BSONType",
"Client",
"CountAggregation",
"CollectionGroup",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
BSONObjectId,
BSONRegex,
BSONTimestamp,
BSONType,
Client,
CollectionGroup,
CollectionReference,
Expand Down Expand Up @@ -108,6 +109,7 @@
"BSONObjectId",
"BSONRegex",
"BSONTimestamp",
"BSONType",
"Client",
"CountAggregation",
"CollectionGroup",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
BSONObjectId,
BSONRegex,
BSONTimestamp,
BSONType,
)
from google.cloud.firestore_v1.client import Client
from google.cloud.firestore_v1.collection import CollectionReference
Expand Down Expand Up @@ -165,6 +166,7 @@
"BSONObjectId",
"BSONRegex",
"BSONTimestamp",
"BSONType",
"Client",
"CountAggregation",
"CollectionGroup",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@
import google
from google.cloud import exceptions # type: ignore
from google.cloud.firestore_v1 import transforms, types
from google.cloud.firestore_v1.bson import _BSONType
from google.cloud.firestore_v1.bson import BSONType
from google.cloud.firestore_v1.field_path import FieldPath, parse_field_path
from google.cloud.firestore_v1.types import common, document, write
from google.cloud.firestore_v1.types.write import DocumentTransform
Expand Down Expand Up @@ -211,7 +211,7 @@ def encode_value(value) -> types.document.Value:
if document_path is not None:
return document.Value(reference_value=document_path)

if isinstance(value, _BSONType):
if isinstance(value, BSONType):
return encode_value(value._to_map_value())

if isinstance(value, GeoPoint):
Expand Down Expand Up @@ -350,7 +350,18 @@ def reference_value_to_document(reference_value, client) -> Any:
def decode_value(
value, client
) -> Union[
None, bool, int, float, list, datetime.datetime, str, bytes, dict, GeoPoint, Vector
None,
bool,
int,
float,
list,
datetime.datetime,
str,
bytes,
dict,
GeoPoint,
Vector,
BSONType,
]:
"""Converts a Firestore protobuf ``Value`` to a native Python value.

Expand All @@ -362,7 +373,9 @@ def decode_value(

Returns:
Union[NoneType, bool, int, float, datetime.datetime, \
str, bytes, dict, ~google.cloud.Firestore.GeoPoint]: A native
str, bytes, dict, ~google.cloud.Firestore.GeoPoint, \
~google.cloud.firestore_v1.vector.Vector, \
~google.cloud.firestore_v1.bson.BSONType]: A native \
Python value converted from the ``value``.

Raises:
Expand Down Expand Up @@ -402,7 +415,10 @@ def decode_value(
raise ValueError("Unknown ``value_type``", value_type)


def decode_dict(value_fields, client) -> Union[dict, Vector]:
def decode_dict(
value_fields,
client,
) -> Union[dict, Vector, BSONType, bytes]:
"""Converts a protobuf map of Firestore ``Value``-s.

Args:
Expand All @@ -412,9 +428,9 @@ def decode_dict(value_fields, client) -> Union[dict, Vector]:
A client that has a document factory.

Returns:
Dict[str, Union[NoneType, bool, int, float, datetime.datetime, \
str, bytes, dict, ~google.cloud.Firestore.GeoPoint]]: A dictionary
of native Python values converted from the ``value_fields``.
Union[dict, ~google.cloud.firestore_v1.vector.Vector, \
~google.cloud.firestore_v1.bson.BSONType, bytes]: A dictionary of native \
Python values, Vector, BSON object, or bytes converted from ``value_fields``.
"""
value_fields_pb = getattr(value_fields, "_pb", value_fields)
res = {key: decode_value(value, client) for key, value in value_fields_pb.items()}
Expand All @@ -425,6 +441,10 @@ def decode_dict(value_fields, client) -> Union[dict, Vector]:
values = cast(Sequence[float], res["value"])
return Vector(values)

decoded = BSONType._from_dict(res)
if decoded is not None:
return decoded

return res


Expand Down
61 changes: 51 additions & 10 deletions packages/google-cloud-firestore/google/cloud/firestore_v1/bson.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,10 @@
import abc
import decimal
import re
from typing import Any, Dict, Union
from typing import Any, Callable, Dict, Optional, Union

__all__ = [
"BSONType",
"BSONObjectId",
"BSONMinKey",
"BSONMaxKey",
Expand All @@ -44,7 +45,7 @@
_HEX_24_REGEX = re.compile(r"^[0-9a-fA-F]{24}$")


class _BSONType(abc.ABC):
class BSONType(abc.ABC):
"""Abstract base class for all BSON type containers in Firestore."""

__slots__ = ()
Expand All @@ -65,11 +66,33 @@ def __eq__(self, other: Any) -> bool:
def __hash__(self) -> int:
"""Hash representation contract for set and dictionary keys."""

@classmethod
def _from_dict(cls, data: Any) -> Optional[Union["BSONType", bytes]]:
"""Deserializes a BSON wire map dictionary into a BSON instance or bytes.

Args:
data (Any): Potential BSON wire map dictionary.

Returns:
Optional[Union[BSONType, bytes]]: Deserialized BSON container
instance or bytes, or None if not a BSON wire map or if decoding fails.
"""
if not isinstance(data, dict) or len(data) != 1:
return None
key, val = next(iter(data.items()))
decoder = _BSON_DECODERS.get(key)
if decoder is None:
return None
try:
return decoder(val)
except Exception:
return None

def __repr__(self) -> str:
return f"{self.__class__.__name__}()"


class BSONObjectId(_BSONType):
class BSONObjectId(BSONType):
"""Represents a 12-byte BSON ObjectId identifier.

Args:
Expand Down Expand Up @@ -128,7 +151,7 @@ def __hash__(self) -> int:
return hash((type(self), self._value))


class BSONMinKey(_BSONType):
class BSONMinKey(BSONType):
"""Represents the BSON MinKey sentinel value for query range boundaries."""

__slots__ = ()
Expand All @@ -146,7 +169,7 @@ def __hash__(self) -> int:
return hash(type(self))


class BSONMaxKey(_BSONType):
class BSONMaxKey(BSONType):
"""Represents the BSON MaxKey sentinel value for query range boundaries."""

__slots__ = ()
Expand All @@ -164,7 +187,7 @@ def __hash__(self) -> int:
return hash(type(self))


class BSONInt32(_BSONType):
class BSONInt32(BSONType):
"""Represents a 32-bit signed integer value container for Firestore BSON.

Args:
Expand Down Expand Up @@ -221,7 +244,7 @@ def __hash__(self) -> int:
return hash((type(self), self._value))


class BSONBinary(_BSONType):
class BSONBinary(BSONType):
"""Represents a BSON binary data container with a subtype for Firestore.

Args:
Expand Down Expand Up @@ -286,7 +309,7 @@ def __hash__(self) -> int:
return hash((type(self), self._data, self._subtype))


class BSONTimestamp(_BSONType):
class BSONTimestamp(BSONType):
"""Container for BSON Timestamp values.

Args:
Expand Down Expand Up @@ -347,7 +370,7 @@ def __hash__(self) -> int:
return hash((type(self), self._seconds, self._increment))


class BSONRegex(_BSONType):
class BSONRegex(BSONType):
"""Represents a BSON Regular Expression container for Firestore.

Args:
Expand Down Expand Up @@ -409,7 +432,7 @@ def __hash__(self) -> int:
return hash((type(self), self._pattern, self._options))


class BSONDecimal128(_BSONType):
class BSONDecimal128(BSONType):
"""Represents a BSON 128-bit Decimal container for Firestore.

Args:
Expand Down Expand Up @@ -507,3 +530,21 @@ def __hash__(self) -> int:
return hash(d)
except decimal.InvalidOperation:
return hash((type(self), self._value))


_BSON_DECODERS: Dict[str, Callable[..., Optional[Union[BSONType, bytes]]]] = {
"__oid__": BSONObjectId,
"__min__": lambda _: BSONMinKey(),
"__max__": lambda _: BSONMaxKey(),
"__int__": BSONInt32,
"__decimal128__": BSONDecimal128,
"__binary__": lambda v: (v[1:] if v[0] == 0 else BSONBinary(v[1:], subtype=v[0]))
if isinstance(v, (bytes, bytearray)) and len(v) >= 1
else None,
"__request_timestamp__": lambda v: BSONTimestamp(v["seconds"], v["increment"])
if isinstance(v, dict) and "seconds" in v and "increment" in v
else None,
"__regex__": lambda v: BSONRegex(v["pattern"], v.get("options", ""))
if isinstance(v, dict) and "pattern" in v
else None,
}
Loading
Loading