diff --git a/docs/changelog.rst b/docs/changelog.rst index 579dfc460..30816cc3b 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -17,6 +17,24 @@ Changes in 1.0.0 - make sure to read https://www.mongodb.com/docs/manual/core/transactions-in-applications/#callback-api-vs-core-api - run_in_transaction context manager relies on Pymongo coreAPI, it will retry automatically in case of ``UnknownTransactionCommitResult`` but not ``TransientTransactionError`` exceptions - Using .count() in a transaction will always use Collection.count_document (as estimated_document_count is not supported in transactions) +- BREAKING CHANGE (#2928): Assigning a field value and deleting a field now + have distinct persistence semantics: + + - Assigning empty default values, such as ``""``, ``[]``, or ``{}``, now + stores those values. Previous versions implicitly unset the field. #267 + - Assigning falsy values to dynamic fields now stores those values instead + of implicitly unsetting the field. + + Existing documents with missing default-valued fields are not migrated + automatically. They must be migrated if queries for the default value need + to match them. + +- BREAKING CHANGE (internal API, supporting #2928): + ``_get_changed_fields()`` was removed and replaced by + ``_get_updated_fields()``, which returns a + ``(changed_fields, unset_fields)`` tuple containing two disjoint lists of + database paths. ``_clear_changed_fields()`` was renamed to + ``_clear_updated_fields()``. - Add a warning that ``mongoengine.org`` is no longer controlled by the MongoEngine project and appears to be an expired domain takeover. - Bug Fix - Fix querying GenericReferenceField with __in operator #2886 diff --git a/mongoengine/base/datastructures.py b/mongoengine/base/datastructures.py index dcb8438c7..d1d7701f7 100644 --- a/mongoengine/base/datastructures.py +++ b/mongoengine/base/datastructures.py @@ -38,6 +38,17 @@ def wrapper(self, key, *args, **kwargs): return wrapper +def mark_key_as_unset_wrapper(parent_method): + """Decorator that ensures _mark_as_unset gets called after deleting a key.""" + + def wrapper(self, key, *args, **kwargs): + result = parent_method(self, key, *args, **kwargs) + self._mark_as_unset(key) + return result + + return wrapper + + class BaseDict(dict): """A special dict so we can watch any changes.""" @@ -86,8 +97,7 @@ def __setstate__(self, state): return self __setitem__ = mark_key_as_changed_wrapper(dict.__setitem__) - __delattr__ = mark_key_as_changed_wrapper(dict.__delattr__) - __delitem__ = mark_key_as_changed_wrapper(dict.__delitem__) + __delitem__ = mark_key_as_unset_wrapper(dict.__delitem__) pop = mark_as_changed_wrapper(dict.pop) clear = mark_as_changed_wrapper(dict.clear) update = mark_as_changed_wrapper(dict.update) @@ -101,6 +111,10 @@ def _mark_as_changed(self, key=None): else: self._instance._mark_as_changed(self._name) + def _mark_as_unset(self, key): + if hasattr(self._instance, "_mark_as_unset"): + self._instance._mark_as_unset(f"{self._name}.{key}") + class BaseList(list): """A special list so we can watch any changes.""" diff --git a/mongoengine/base/document.py b/mongoengine/base/document.py index 366915a8e..e2d6ac43c 100644 --- a/mongoengine/base/document.py +++ b/mongoengine/base/document.py @@ -1,5 +1,4 @@ import copy -import numbers import warnings from functools import partial @@ -12,7 +11,6 @@ BaseDict, BaseList, EmbeddedDocumentList, - LazyReference, StrictDict, ) from mongoengine.base.fields import ComplexBaseField @@ -32,18 +30,10 @@ class BaseDocument: - # TODO simplify how `_changed_fields` is used. - # Currently, handling of `_changed_fields` seems unnecessarily convoluted: - # 1. `BaseDocument` defines `_changed_fields` in its `__slots__`, yet it's - # not setting it to `[]` (or any other value) in `__init__`. - # 2. `EmbeddedDocument` sets `_changed_fields` to `[]` it its overloaded - # `__init__`. - # 3. `Document` does NOT set `_changed_fields` upon initialization. The - # field is primarily set via `_from_son` or `_clear_changed_fields`, - # though there are also other methods that manipulate it. - # 4. The codebase is littered with `hasattr` calls for `_changed_fields`. __slots__ = ( "_changed_fields", + "_unset_fields", + "_has_change_tracking_baseline", "_initialised", "_created", "_data", @@ -67,9 +57,18 @@ def __init__(self, *args, **values): to Python-type values via each field's `to_python` method. :param _created: Indicates whether this is a brand new document or whether it's already been persisted before. Defaults to true. + :param _changed_fields: Database field paths changed since the tracking + baseline was established. + :param _unset_fields: Database field paths explicitly unset since the + tracking baseline was established. + :param _has_change_tracking_baseline: Whether the document has a known + state against which changes can be tracked. """ self._initialised = False self._created = True + self._changed_fields = [] + self._unset_fields = [] + self._has_change_tracking_baseline = False if args: raise TypeError( @@ -150,10 +149,17 @@ def __delattr__(self, *args, **kwargs): if callable(default): default = default() setattr(self, field_name, default) + # Without a baseline, preserve historical full-document behavior + # and let serialization determine the delta. + if self._has_change_tracking_baseline: + self._mark_as_unset(field_name) else: super().__delattr__(*args, **kwargs) def __setattr__(self, name, value): + if not name.startswith("_") and self._unset_fields: + self._unmark_as_unset(name) + # Handle dynamic data only if an initialised dynamic document if self._dynamic and not self._dynamic_lock: if name not in self._fields_ordered and not name.startswith("_"): @@ -169,7 +175,7 @@ def __setattr__(self, name, value): # Handle marking data as changed if name in self._dynamic_fields: self._data[name] = value - if hasattr(self, "_changed_fields"): + if self._initialised: self._mark_as_changed(name) try: self__created = self._created @@ -207,6 +213,8 @@ def __getstate__(self): data = {} for k in ( "_changed_fields", + "_unset_fields", + "_has_change_tracking_baseline", "_initialised", "_created", "_dynamic_fields", @@ -219,10 +227,19 @@ def __getstate__(self): return data def __setstate__(self, data): + # __init__ is not called when unpickling, and legacy states may omit + # these tracking fields so it's important that we set them to maintain backward compatibility. + self._changed_fields = [] + self._unset_fields = [] + self._has_change_tracking_baseline = ( + not self._is_document or "_changed_fields" in data + ) if data.pop("_data_is_mongo", False) or isinstance(data["_data"], SON): data["_data"] = self.__class__._from_son(data["_data"])._data for k in ( "_changed_fields", + "_unset_fields", + "_has_change_tracking_baseline", "_initialised", "_created", "_data", @@ -517,17 +534,50 @@ def __expand_dynamic_values(self, name, value): return value + def _resolve_key(self, key): + """Resolve a field name to its database path.""" + if "." in key: + key, rest = key.split(".", 1) + key = self._db_field_map.get(key, key) + return f"{key}.{rest}" + return self._db_field_map.get(key, key) + + def _unmark_as_unset(self, key): + if not key or not self._unset_fields: + return + + key = self._resolve_key(key) + self._unset_fields = [ + path + for path in self._unset_fields + if not ( + path == key or path.startswith(f"{key}.") or key.startswith(f"{path}.") + ) + ] + + def _mark_as_unset(self, key): + if not key: + return + + key = self._resolve_key(key) + # An exact or ancestor unset already removes this path. + if key in self._unset_fields or any( + key.startswith(f"{path}.") for path in self._unset_fields + ): + return + # Unsetting an ancestor supersedes previously tracked descendant unsets. + self._unset_fields = [ + path for path in self._unset_fields if not path.startswith(f"{key}.") + ] + self._unset_fields.append(key) + def _mark_as_changed(self, key): """Mark a key as explicitly changed by the user.""" - if not hasattr(self, "_changed_fields"): + if not self._has_change_tracking_baseline: return - if "." in key: - key, rest = key.split(".", 1) - key = self._db_field_map.get(key, key) - key = f"{key}.{rest}" - else: - key = self._db_field_map.get(key, key) + self._unmark_as_unset(key) + key = self._resolve_key(key) if key not in self._changed_fields: levels, idx = key.split("."), 1 @@ -544,15 +594,16 @@ def _mark_as_changed(self, key): if field.startswith(level): remove(field) - def _clear_changed_fields(self): - """Using _get_changed_fields iterate and remove any fields that + def _clear_updated_fields(self): + """Using _get_updated_fields iterate and remove any fields that are marked as changed. """ ReferenceField = _import_class("ReferenceField") GenericReferenceField = _import_class("GenericReferenceField") - for changed in self._get_changed_fields(): - parts = changed.split(".") + changed_fields, unset_fields = self._get_updated_fields() + for updated_field in changed_fields + unset_fields: + parts = updated_field.split(".") data = self for part in parts: if isinstance(data, list): @@ -566,25 +617,26 @@ def _clear_changed_fields(self): field_name = data._reverse_db_field_map.get(part, part) data = getattr(data, field_name, None) - if not isinstance(data, LazyReference) and hasattr( - data, "_changed_fields" - ): + if isinstance(data, BaseDocument): if getattr(data, "_is_document", False): continue data._changed_fields = [] + data._unset_fields = [] elif isinstance(data, (list, tuple, dict)): if hasattr(data, "field") and isinstance( data.field, (ReferenceField, GenericReferenceField) ): continue - BaseDocument._nestable_types_clear_changed_fields(data) + BaseDocument._nestable_types_clear_updated_fields(data) self._changed_fields = [] + self._unset_fields = [] + self._has_change_tracking_baseline = True @staticmethod - def _nestable_types_clear_changed_fields(data): - """Inspect nested data for changed fields + def _nestable_types_clear_updated_fields(data): + """Inspect nested data for updated fields :param data: data to inspect for changes """ @@ -598,18 +650,19 @@ def _nestable_types_clear_changed_fields(data): iterator = data.items() for _index_or_key, value in iterator: - if hasattr(value, "_get_changed_fields") and not isinstance( + if hasattr(value, "_get_updated_fields") and not isinstance( value, Document ): # don't follow references - value._clear_changed_fields() + value._clear_updated_fields() elif isinstance(value, (list, tuple, dict)): - BaseDocument._nestable_types_clear_changed_fields(value) + BaseDocument._nestable_types_clear_updated_fields(value) @staticmethod - def _nestable_types_changed_fields(changed_fields, base_key, data): - """Inspect nested data for changed fields + def _nestable_types_updated_fields(changed_fields, unset_fields, base_key, data): + """Inspect nested data for updated fields :param changed_fields: Previously collected changed fields + :param unset_fields: Previously collected unset fields :param base_key: The base key that must be used to prepend changes to this data :param data: data to inspect for changes """ @@ -623,20 +676,39 @@ def _nestable_types_changed_fields(changed_fields, base_key, data): for index_or_key, value in iterator: item_key = f"{base_key}{index_or_key}." # don't check anything lower if this key is already marked - # as changed. - if item_key[:-1] in changed_fields: + # as changed or unset. + if item_key[:-1] in changed_fields or item_key[:-1] in unset_fields: continue - if hasattr(value, "_get_changed_fields"): - changed = value._get_changed_fields() + if hasattr(value, "_get_updated_fields"): + changed, unset = value._get_updated_fields() changed_fields += [f"{item_key}{k}" for k in changed if k] + unset_fields += [f"{item_key}{k}" for k in unset if k] elif isinstance(value, (list, tuple, dict)): - BaseDocument._nestable_types_changed_fields( - changed_fields, item_key, value + BaseDocument._nestable_types_updated_fields( + changed_fields, unset_fields, item_key, value ) - def _get_changed_fields(self): - """Return a list of all fields that have explicitly been changed.""" + @staticmethod + def _get_disjoint_updated_fields(changed_fields, unset_fields): + """Remove ancestor and descendant conflicts from updated field paths.""" + unset_fields = [ + field + for field in unset_fields + if not any(field.startswith(f"{changed}.") for changed in changed_fields) + ] + changed_fields = [ + field + for field in changed_fields + if not any( + field == unset or field.startswith(f"{unset}.") + for unset in unset_fields + ) + ] + return changed_fields, unset_fields + + def _get_updated_fields(self): + """Return lists of fields that were explicitly changed or unset.""" EmbeddedDocument = _import_class("EmbeddedDocument") LazyReferenceField = _import_class("LazyReferenceField") ReferenceField = _import_class("ReferenceField") @@ -644,8 +716,8 @@ def _get_changed_fields(self): GenericReferenceField = _import_class("GenericReferenceField") SortedListField = _import_class("SortedListField") - changed_fields = [] - changed_fields += getattr(self, "_changed_fields", []) + changed_fields = list(self._changed_fields) + unset_fields = list(self._unset_fields) for field_name in self._fields_ordered: db_field_name = self._db_field_map.get(field_name, field_name) @@ -653,18 +725,24 @@ def _get_changed_fields(self): data = self._data.get(field_name, None) field = self._fields.get(field_name) - if db_field_name in changed_fields: - # Whole field already marked as changed, no need to go further - continue - if isinstance(field, ReferenceField): # Don't follow referenced documents continue if isinstance(data, EmbeddedDocument): - # Find all embedded fields that have been changed - changed = data._get_changed_fields() + changed, unset = data._get_updated_fields() + if db_field_name in unset_fields: + if changed: + unset_fields.remove(db_field_name) + if db_field_name not in changed_fields: + changed_fields.append(db_field_name) + continue + if db_field_name in changed_fields: + continue changed_fields += [f"{key}{k}" for k in changed if k] + unset_fields += [f"{key}{k}" for k in unset if k] elif isinstance(data, (list, tuple, dict)): + if db_field_name in changed_fields or db_field_name in unset_fields: + continue if hasattr(field, "field") and isinstance( field.field, ( @@ -681,8 +759,25 @@ def _get_changed_fields(self): changed_fields.append(db_field_name) continue - self._nestable_types_changed_fields(changed_fields, key, data) - return changed_fields + self._nestable_types_updated_fields( + changed_fields, unset_fields, key, data + ) + + return self._get_disjoint_updated_fields(changed_fields, unset_fields) + + @staticmethod + def _remove_path(data, path): + """Remove a database path from serialized document data.""" + parts = path.split(".") + for part in parts[:-1]: + if isinstance(data, list) and part.isdigit(): + data = data[int(part)] + elif hasattr(data, "get"): + data = data.get(part) + if data is None: + break + else: + data.pop(parts[-1], None) def _delta(self): """Returns the delta (set, unset) of the changes for a document. @@ -691,80 +786,46 @@ def _delta(self): # Handles cases where not loaded from_son but has _id doc = self.to_mongo() - set_fields = self._get_changed_fields() - unset_data = {} - if hasattr(self, "_changed_fields"): + set_fields, unset_fields = self._get_updated_fields() + if self._has_change_tracking_baseline: set_data = {} + unset_data = {path: 1 for path in unset_fields} # Fetch each set item from its path for path in set_fields: parts = path.split(".") d = doc new_path = [] + missing = False for p in parts: if isinstance(d, (ObjectId, DBRef)): # Don't dig in the references break - elif isinstance(d, list) and p.isdigit(): + + new_path.append(p) + if missing: + continue + if isinstance(d, list) and p.isdigit(): # An item of a list (identified by its index) is updated d = d[int(p)] elif hasattr(d, "get"): # dict-like (dict, embedded document) - d = d.get(p) - new_path.append(p) + if p in d: + d = d[p] + else: + missing = True path = ".".join(new_path) - set_data[path] = d + if missing: + unset_data[path] = 1 + else: + set_data[path] = d else: set_data = doc + unset_data = {} if "_id" in set_data: del set_data["_id"] + for path in unset_fields: + self._remove_path(set_data, path) - # Determine if any changed items were actually unset. - for path, value in list(set_data.items()): - if value or isinstance( - value, (numbers.Number, bool) - ): # Account for 0 and True that are truthy - continue - - parts = path.split(".") - - if self._dynamic and len(parts) and parts[0] in self._dynamic_fields: - del set_data[path] - unset_data[path] = 1 - continue - - # If we've set a value that ain't the default value don't unset it. - default = None - if path in self._fields: - default = self._fields[path].default - else: # Perform a full lookup for lists / embedded lookups - d = self - db_field_name = parts.pop() - for p in parts: - if isinstance(d, list) and p.isdigit(): - d = d[int(p)] - elif hasattr(d, "__getattribute__") and not isinstance(d, dict): - real_path = d._reverse_db_field_map.get(p, p) - d = getattr(d, real_path) - else: - d = d.get(p) - - if hasattr(d, "_fields"): - field_name = d._reverse_db_field_map.get( - db_field_name, db_field_name - ) - if field_name in d._fields: - default = d._fields.get(field_name).default - else: - default = None - - if default is not None: - default = default() if callable(default) else default - - if value != default: - continue - - del set_data[path] - unset_data[path] = 1 return set_data, unset_data @classmethod @@ -838,7 +899,7 @@ def _from_son(cls, son, _auto_dereference=True, created=False): data = {k: v for k, v in data.items() if k in cls._fields} obj = cls(__auto_convert=False, _created=created, **data) - obj._changed_fields = [] + obj._has_change_tracking_baseline = True if not _auto_dereference: obj._fields = fields diff --git a/mongoengine/document.py b/mongoengine/document.py index ef2f5173f..01a55cdc3 100644 --- a/mongoengine/document.py +++ b/mongoengine/document.py @@ -94,7 +94,7 @@ class EmbeddedDocument(BaseDocument, metaclass=DocumentMetaclass): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._instance = None - self._changed_fields = [] + self._has_change_tracking_baseline = True def __eq__(self, other): if isinstance(other, self.__class__): @@ -353,6 +353,8 @@ def modify(self, query=None, **update): setattr(self, field, self._reload(field, updated[field])) self._changed_fields = updated._changed_fields + self._unset_fields = updated._unset_fields + self._has_change_tracking_baseline = updated._has_change_tracking_baseline self._created = False return True @@ -494,7 +496,7 @@ def save( self.__class__, document=self, created=created, **signal_kwargs ) - self._clear_changed_fields() + self._clear_updated_fields() self._created = False return self @@ -612,7 +614,9 @@ def cascade_save(self, **kwargs): if not ref or isinstance(ref, DBRef): continue - if not getattr(ref, "_changed_fields", True): + if ref._has_change_tracking_baseline and not ( + ref._changed_fields or ref._unset_fields + ): continue ref_id = f"{ref.__class__.__name__},{str(ref._data)}" @@ -620,7 +624,6 @@ def cascade_save(self, **kwargs): _refs.append(ref_id) kwargs["_refs"] = _refs ref.save(**kwargs) - ref._changed_fields = [] @property def _qs(self): @@ -816,6 +819,10 @@ def reload(self, *fields, **kwargs): if fields else obj._changed_fields ) + self._unset_fields = ( + list(set(self._unset_fields) - set(fields)) if fields else obj._unset_fields + ) + self._has_change_tracking_baseline = True self._created = False return self @@ -835,6 +842,7 @@ def _reload(self, key, value): elif isinstance(value, (EmbeddedDocument, DynamicEmbeddedDocument)): value._instance = None value._changed_fields = [] + value._unset_fields = [] return value def to_dbref(self): @@ -1087,13 +1095,11 @@ class DynamicDocument(Document, metaclass=TopLevelDocumentMetaclass): _dynamic = True def __delattr__(self, *args, **kwargs): - """Delete the attribute by setting to None and allowing _delta - to unset it. - """ + """Delete a dynamic field.""" field_name = args[0] if field_name in self._dynamic_fields: setattr(self, field_name, None) - self._dynamic_fields[field_name].null = False + self._mark_as_unset(field_name) else: super().__delattr__(*args, **kwargs) @@ -1110,17 +1116,13 @@ class DynamicEmbeddedDocument(EmbeddedDocument, metaclass=DocumentMetaclass): _dynamic = True def __delattr__(self, *args, **kwargs): - """Delete the attribute by setting to None and allowing _delta - to unset it. - """ + """Delete a dynamic embedded field.""" field_name = args[0] if field_name in self._fields: - default = self._fields[field_name].default - if callable(default): - default = default() - setattr(self, field_name, default) + super().__delattr__(*args, **kwargs) else: setattr(self, field_name, None) + self._mark_as_unset(field_name) class MapReduceDocument: diff --git a/tests/document/test_delta.py b/tests/document/test_delta.py index e4d4fa7bd..e91225dae 100644 --- a/tests/document/test_delta.py +++ b/tests/document/test_delta.py @@ -1,5 +1,7 @@ import unittest +import pytest + from mongoengine import * from mongoengine.pymongo_support import list_collection_names from tests.utils import MongoDBTestCase, get_as_pymongo @@ -27,6 +29,317 @@ def test_delta(self): self.delta(Document) self.delta(DynamicDocument) + def test_updated_fields__new_documents__initializes_empty_tracking_lists(self): + class Embedded(EmbeddedDocument): + value = StringField() + + class Doc(Document): + value = StringField() + + class Dynamic(DynamicDocument): + pass + + for document in ( + Embedded(value="value"), + Doc(value="value"), + Dynamic(value="value"), + ): + assert document._changed_fields == [] + assert document._unset_fields == [] + + def test_delta__primary_key_assigned_to_new_document__uses_full_document(self): + class Doc(Document): + key = StringField(primary_key=True) + value = StringField() + + doc = Doc(value="value") + doc.key = "key" + + assert doc._created is False + # key is MongoDB's _id, which _delta() deliberately excludes. + assert doc._delta() == ({"value": "value"}, {}) + + def test_save__post_init_primary_key_and_deleted_field__keeps_historical_noop( + self, + ): + class Doc(Document): + key = StringField(primary_key=True) + value = StringField() + + doc = Doc(value="value") + doc.key = "key" + del doc.value + + assert doc._delta() == ({}, {}) + + doc.save() + + # This silent no-op is a historical bug. Fixing it is worth considering + # separately because persisting this document would be a breaking change. + assert Doc.objects(pk=doc.pk).count() == 0 + + def test_save__post_init_primary_key_and_deleted_default__preserves_default(self): + class Doc(Document): + key = StringField(primary_key=True) + value = IntField(default=4) + + doc = Doc() + doc.key = "key" + del doc.value + + assert doc.value == 4 + assert doc._delta() == ({"value": 4}, {}) + + doc.save() + + assert Doc._get_collection().find_one({"_id": doc.pk}) == { + "_id": doc.pk, + "value": 4, + } + + def test_delta__default_values_are_assigned__sets_values(self): + class Doc(Document): + empty_list = ListField() + empty_dict = DictField() + empty_string = StringField(default="") + false_value = BooleanField(default=False) + zero_value = IntField(default=0) + callable_value = StringField(default=lambda: "default", db_field="value") + + doc = Doc( + empty_list=[1], + empty_dict={"key": "value"}, + empty_string="value", + false_value=True, + zero_value=1, + callable_value="value", + ).save() + + doc.empty_list = [] + doc.empty_dict = {} + doc.empty_string = "" + doc.false_value = False + doc.zero_value = 0 + doc.callable_value = "default" + + assert doc._delta() == ( + { + "empty_list": [], + "empty_dict": {}, + "empty_string": "", + "false_value": False, + "zero_value": 0, + "value": "default", + }, + {}, + ) + + doc.save() + assert get_as_pymongo(doc) == { + "_id": doc.id, + "empty_list": [], + "empty_dict": {}, + "empty_string": "", + "false_value": False, + "zero_value": 0, + "value": "default", + } + + def test_save__default_values_are_assigned__stores_queryable_values(self): + class Doc(Document): + items = ListField() + enabled = BooleanField(default=False) + + doc = Doc(items=[1], enabled=True).save() + doc.items = [] + doc.enabled = False + doc.save() + + assert get_as_pymongo(doc) == { + "_id": doc.id, + "items": [], + "enabled": False, + } + assert Doc.objects(items=[], enabled=False).count() == 1 + + def test_delta__field_is_deleted_then_assigned_default__sets_default(self): + class Doc(Document): + values = ListField() + + doc = Doc(values=[1]).save() + + del doc.values + doc.values = [] + + assert doc._delta() == ({"values": []}, {}) + + def test_delta__deleted_list_is_mutated__sets_list(self): + class Doc(Document): + values = ListField(db_field="db_values") + db_values = StringField(db_field="other_value") + + doc = Doc(values=[1], db_values="preserved").save() + + del doc.values + doc.values.append(2) + + assert doc._delta() == ({"db_values": [2]}, {}) + + def test_delta__deleted_embedded_default_is_modified__sets_whole_embedded(self): + class Settings(EmbeddedDocument): + level = IntField(default=4) + + class Account(Document): + settings = EmbeddedDocumentField(Settings, default=Settings) + + account = Account(settings=Settings(level=10)).save() + + del account.settings + + assert account.settings.level == 4 + assert account._get_updated_fields() == ([], ["settings"]) + assert account._delta() == ({}, {"settings": 1}) + + account.settings.level = 7 + + assert account._get_updated_fields() == (["settings"], []) + assert account._delta() == ({"settings": {"level": 7}}, {}) + + account.save() + assert get_as_pymongo(account) == { + "_id": account.id, + "settings": {"level": 7}, + } + + @pytest.mark.xfail( + reason=( + "A child unset is lost when its deleted embedded parent is restored " + "and promoted to a whole-field set" + ), + strict=True, + ) + def test_save__deleted_embedded_default_is_modified_and_child_deleted__omits_deleted_child( + self, + ): + class Settings(EmbeddedDocument): + theme = StringField(default="system") + language = StringField(default="en") + + class Account(Document): + settings = EmbeddedDocumentField(Settings, default=Settings) + + account = Account(settings=Settings(theme="dark", language="fr")).save() + + del account.settings + account.settings.language = "de" + del account.settings.theme + + assert account.settings._get_updated_fields() == ( + ["language"], + ["theme"], + ) + + delta = account._delta() + account.save() + + assert delta == ({"settings": {"language": "de"}}, {}) + assert get_as_pymongo(account) == { + "_id": account.id, + "settings": {"language": "de"}, + } + + def test_get_updated_fields__dict_key_is_deleted__reports_unset(self): + class Doc(Document): + mapping = DictField() + + doc = Doc(mapping={"key": "value"}).save() + + del doc.mapping["key"] + + assert doc._get_updated_fields() == ([], ["mapping.key"]) + assert doc._delta() == ({}, {"mapping.key": 1}) + + def test_delta__dict_is_assigned_then_key_is_deleted__sets_whole_dict(self): + class Doc(Document): + mapping = DictField() + + doc = Doc(mapping={"old": "value"}).save() + + doc.mapping = {"kept": "value", "removed": "value"} + del doc.mapping["removed"] + + assert doc._get_updated_fields() == (["mapping"], []) + assert doc._delta() == ({"mapping": {"kept": "value"}}, {}) + + doc.save() + assert get_as_pymongo(doc) == { + "_id": doc.id, + "mapping": {"kept": "value"}, + } + + def test_delta__dict_key_is_modified_then_field_is_deleted__unsets_dict(self): + class Doc(Document): + mapping = DictField() + + doc = Doc(mapping={"key": "value"}).save() + + doc.mapping["key"] = "changed" + del doc.mapping + + assert doc._get_updated_fields() == ([], ["mapping"]) + assert doc._delta() == ({}, {"mapping": 1}) + + doc.save() + assert get_as_pymongo(doc) == {"_id": doc.id} + + def test_delta__nested_field_is_deleted_on_new_document__omits_field(self): + class Settings(EmbeddedDocument): + number = IntField(default=4) + + class Doc(Document): + settings = EmbeddedDocumentField(Settings, default=Settings) + + doc = Doc() + del doc.settings.number + + assert doc.settings.number == 4 + assert doc._get_updated_fields() == ([], ["settings.number"]) + assert doc._delta() == ({"settings": {}}, {}) + + def test_delta__field_with_truthy_default_is_deleted__unsets_field(self): + class Doc(Document): + value = StringField(default="default", db_field="db_value") + db_value = StringField(db_field="other_value") + + doc = Doc(value="other", db_value="preserved").save() + + del doc.value + + assert doc.value == "default" + assert doc._get_updated_fields() == ([], ["db_value"]) + assert doc._delta() == ({}, {"db_value": 1}) + + doc.save() + + assert get_as_pymongo(doc) == { + "_id": doc.id, + "other_value": "preserved", + } + + doc.reload() + assert doc.value == "default" + + def test_delta__dynamic_field_is_deleted_then_assigned_none__sets_none(self): + class Doc(DynamicDocument): + pass + + doc = Doc(value="other").save() + + del doc.value + doc.value = None + + assert doc._delta() == ({"value": None}, {}) + @staticmethod def delta(DocClass): class Doc(DocClass): @@ -40,40 +353,50 @@ class Doc(DocClass): doc.save() doc = Doc.objects.first() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert doc._delta() == ({}, {}) doc.string_field = "hello" - assert doc._get_changed_fields() == ["string_field"] + assert doc._get_updated_fields() == (["string_field"], []) assert doc._delta() == ({"string_field": "hello"}, {}) doc._changed_fields = [] doc.int_field = 1 - assert doc._get_changed_fields() == ["int_field"] + assert doc._get_updated_fields() == (["int_field"], []) assert doc._delta() == ({"int_field": 1}, {}) doc._changed_fields = [] dict_value = {"hello": "world", "ping": "pong"} doc.dict_field = dict_value - assert doc._get_changed_fields() == ["dict_field"] + assert doc._get_updated_fields() == (["dict_field"], []) assert doc._delta() == ({"dict_field": dict_value}, {}) doc._changed_fields = [] list_value = ["1", 2, {"hello": "world"}] doc.list_field = list_value - assert doc._get_changed_fields() == ["list_field"] + assert doc._get_updated_fields() == (["list_field"], []) assert doc._delta() == ({"list_field": list_value}, {}) - # Test unsetting + # Test assigning empty defaults doc._changed_fields = [] doc.dict_field = {} - assert doc._get_changed_fields() == ["dict_field"] - assert doc._delta() == ({}, {"dict_field": 1}) + assert doc._get_updated_fields() == (["dict_field"], []) + assert doc._delta() == ({"dict_field": {}}, {}) doc._changed_fields = [] doc.list_field = [] - assert doc._get_changed_fields() == ["list_field"] - assert doc._delta() == ({}, {"list_field": 1}) + assert doc._get_updated_fields() == (["list_field"], []) + assert doc._delta() == ({"list_field": []}, {}) + + # Test explicit unsetting + for field_name in ("int_field", "string_field", "list_field", "dict_field"): + doc._changed_fields = [] + doc._unset_fields = [] + + delattr(doc, field_name) + + assert doc._get_updated_fields() == ([], [field_name]) + assert doc._delta() == ({}, {field_name: 1}) def test_delta_recursive(self): self.delta_recursive(Document, EmbeddedDocument) @@ -101,7 +424,7 @@ class Doc(DocClass): doc.save() doc = Doc.objects.first() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert doc._delta() == ({}, {}) embedded_1 = Embedded() @@ -112,7 +435,7 @@ class Doc(DocClass): embedded_1.list_field = ["1", 2, {"hello": "world"}] doc.embedded_field = embedded_1 - assert doc._get_changed_fields() == ["embedded_field"] + assert doc._get_updated_fields() == (["embedded_field"], []) embedded_delta = { "id": "010101", @@ -128,17 +451,17 @@ class Doc(DocClass): doc = doc.reload(10) doc.embedded_field.dict_field = {} - assert doc._get_changed_fields() == ["embedded_field.dict_field"] - assert doc.embedded_field._delta() == ({}, {"dict_field": 1}) - assert doc._delta() == ({}, {"embedded_field.dict_field": 1}) + assert doc._get_updated_fields() == (["embedded_field.dict_field"], []) + assert doc.embedded_field._delta() == ({"dict_field": {}}, {}) + assert doc._delta() == ({"embedded_field.dict_field": {}}, {}) doc.save() doc = doc.reload(10) assert doc.embedded_field.dict_field == {} doc.embedded_field.list_field = [] - assert doc._get_changed_fields() == ["embedded_field.list_field"] - assert doc.embedded_field._delta() == ({}, {"list_field": 1}) - assert doc._delta() == ({}, {"embedded_field.list_field": 1}) + assert doc._get_updated_fields() == (["embedded_field.list_field"], []) + assert doc.embedded_field._delta() == ({"list_field": []}, {}) + assert doc._delta() == ({"embedded_field.list_field": []}, {}) doc.save() doc = doc.reload(10) assert doc.embedded_field.list_field == [] @@ -150,7 +473,7 @@ class Doc(DocClass): embedded_2.list_field = ["1", 2, {"hello": "world"}] doc.embedded_field.list_field = ["1", 2, embedded_2] - assert doc._get_changed_fields() == ["embedded_field.list_field"] + assert doc._get_updated_fields() == (["embedded_field.list_field"], []) assert doc.embedded_field._delta() == ( { @@ -194,7 +517,10 @@ class Doc(DocClass): assert doc.embedded_field.list_field[2][k] == embedded_2[k] doc.embedded_field.list_field[2].string_field = "world" - assert doc._get_changed_fields() == ["embedded_field.list_field.2.string_field"] + assert doc._get_updated_fields() == ( + ["embedded_field.list_field.2.string_field"], + [], + ) assert doc.embedded_field._delta() == ( {"list_field.2.string_field": "world"}, {}, @@ -210,7 +536,7 @@ class Doc(DocClass): # Test multiple assignments doc.embedded_field.list_field[2].string_field = "hello world" doc.embedded_field.list_field[2] = doc.embedded_field.list_field[2] - assert doc._get_changed_fields() == ["embedded_field.list_field.2"] + assert doc._get_updated_fields() == (["embedded_field.list_field.2"], []) assert doc.embedded_field._delta() == ( { "list_field.2": { @@ -281,7 +607,7 @@ class Doc(DocClass): doc = doc.reload(10) doc.dict_field["Embedded"].string_field = "Hello World" - assert doc._get_changed_fields() == ["dict_field.Embedded.string_field"] + assert doc._get_updated_fields() == (["dict_field.Embedded.string_field"], []) assert doc._delta() == ({"dict_field.Embedded.string_field": "Hello World"}, {}) def test_circular_reference_deltas(self): @@ -376,40 +702,51 @@ class Doc(DocClass): doc.save() doc = Doc.objects.first() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert doc._delta() == ({}, {}) doc.string_field = "hello" - assert doc._get_changed_fields() == ["db_string_field"] + assert doc._get_updated_fields() == (["db_string_field"], []) assert doc._delta() == ({"db_string_field": "hello"}, {}) doc._changed_fields = [] doc.int_field = 1 - assert doc._get_changed_fields() == ["db_int_field"] + assert doc._get_updated_fields() == (["db_int_field"], []) assert doc._delta() == ({"db_int_field": 1}, {}) doc._changed_fields = [] dict_value = {"hello": "world", "ping": "pong"} doc.dict_field = dict_value - assert doc._get_changed_fields() == ["db_dict_field"] + assert doc._get_updated_fields() == (["db_dict_field"], []) assert doc._delta() == ({"db_dict_field": dict_value}, {}) doc._changed_fields = [] list_value = ["1", 2, {"hello": "world"}] doc.list_field = list_value - assert doc._get_changed_fields() == ["db_list_field"] + assert doc._get_updated_fields() == (["db_list_field"], []) assert doc._delta() == ({"db_list_field": list_value}, {}) - # Test unsetting + # Test assigning empty defaults doc._changed_fields = [] doc.dict_field = {} - assert doc._get_changed_fields() == ["db_dict_field"] - assert doc._delta() == ({}, {"db_dict_field": 1}) + assert doc._get_updated_fields() == (["db_dict_field"], []) + assert doc._delta() == ({"db_dict_field": {}}, {}) doc._changed_fields = [] doc.list_field = [] - assert doc._get_changed_fields() == ["db_list_field"] - assert doc._delta() == ({}, {"db_list_field": 1}) + assert doc._get_updated_fields() == (["db_list_field"], []) + assert doc._delta() == ({"db_list_field": []}, {}) + + # Test explicit unsetting + for field_name in ("int_field", "string_field", "list_field", "dict_field"): + db_field_name = f"db_{field_name}" + doc._changed_fields = [] + doc._unset_fields = [] + + delattr(doc, field_name) + + assert doc._get_updated_fields() == ([], [db_field_name]) + assert doc._delta() == ({}, {db_field_name: 1}) # Test it saves that data doc = Doc() @@ -461,7 +798,7 @@ class Doc(DocClass): doc.save() doc = Doc.objects.first() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert doc._delta() == ({}, {}) embedded_1 = Embedded() @@ -471,7 +808,7 @@ class Doc(DocClass): embedded_1.list_field = ["1", 2, {"hello": "world"}] doc.embedded_field = embedded_1 - assert doc._get_changed_fields() == ["db_embedded_field"] + assert doc._get_updated_fields() == (["db_embedded_field"], []) embedded_delta = { "db_string_field": "hello", @@ -486,18 +823,18 @@ class Doc(DocClass): doc = doc.reload(10) doc.embedded_field.dict_field = {} - assert doc._get_changed_fields() == ["db_embedded_field.db_dict_field"] - assert doc.embedded_field._delta() == ({}, {"db_dict_field": 1}) - assert doc._delta() == ({}, {"db_embedded_field.db_dict_field": 1}) + assert doc._get_updated_fields() == (["db_embedded_field.db_dict_field"], []) + assert doc.embedded_field._delta() == ({"db_dict_field": {}}, {}) + assert doc._delta() == ({"db_embedded_field.db_dict_field": {}}, {}) doc.save() doc = doc.reload(10) assert doc.embedded_field.dict_field == {} - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) doc.embedded_field.list_field = [] - assert doc._get_changed_fields() == ["db_embedded_field.db_list_field"] - assert doc.embedded_field._delta() == ({}, {"db_list_field": 1}) - assert doc._delta() == ({}, {"db_embedded_field.db_list_field": 1}) + assert doc._get_updated_fields() == (["db_embedded_field.db_list_field"], []) + assert doc.embedded_field._delta() == ({"db_list_field": []}, {}) + assert doc._delta() == ({"db_embedded_field.db_list_field": []}, {}) doc.save() doc = doc.reload(10) assert doc.embedded_field.list_field == [] @@ -509,7 +846,7 @@ class Doc(DocClass): embedded_2.list_field = ["1", 2, {"hello": "world"}] doc.embedded_field.list_field = ["1", 2, embedded_2] - assert doc._get_changed_fields() == ["db_embedded_field.db_list_field"] + assert doc._get_updated_fields() == (["db_embedded_field.db_list_field"], []) assert doc.embedded_field._delta() == ( { "db_list_field": [ @@ -544,7 +881,7 @@ class Doc(DocClass): {}, ) doc.save() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) doc = doc.reload(10) assert doc.embedded_field.list_field[0] == "1" @@ -553,9 +890,10 @@ class Doc(DocClass): assert doc.embedded_field.list_field[2][k] == embedded_2[k] doc.embedded_field.list_field[2].string_field = "world" - assert doc._get_changed_fields() == [ - "db_embedded_field.db_list_field.2.db_string_field" - ] + assert doc._get_updated_fields() == ( + ["db_embedded_field.db_list_field.2.db_string_field"], + [], + ) assert doc.embedded_field._delta() == ( {"db_list_field.2.db_string_field": "world"}, {}, @@ -571,7 +909,7 @@ class Doc(DocClass): # Test multiple assignments doc.embedded_field.list_field[2].string_field = "hello world" doc.embedded_field.list_field[2] = doc.embedded_field.list_field[2] - assert doc._get_changed_fields() == ["db_embedded_field.db_list_field.2"] + assert doc._get_updated_fields() == (["db_embedded_field.db_list_field.2"], []) assert doc.embedded_field._delta() == ( { "db_list_field.2": { @@ -679,13 +1017,13 @@ class Person(DynamicDocument): p.age = 24 assert p.age == 24 - assert p._get_changed_fields() == ["age"] + assert p._get_updated_fields() == (["age"], []) assert p._delta() == ({"age": 24}, {}) p = Person.objects(age=22).get() p.age = 24 assert p.age == 24 - assert p._get_changed_fields() == ["age"] + assert p._get_updated_fields() == (["age"], []) assert p._delta() == ({"age": 24}, {}) p.save() @@ -700,40 +1038,40 @@ class Doc(DynamicDocument): doc.save() doc = Doc.objects.first() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert doc._delta() == ({}, {}) doc.string_field = "hello" - assert doc._get_changed_fields() == ["string_field"] + assert doc._get_updated_fields() == (["string_field"], []) assert doc._delta() == ({"string_field": "hello"}, {}) doc._changed_fields = [] doc.int_field = 1 - assert doc._get_changed_fields() == ["int_field"] + assert doc._get_updated_fields() == (["int_field"], []) assert doc._delta() == ({"int_field": 1}, {}) doc._changed_fields = [] dict_value = {"hello": "world", "ping": "pong"} doc.dict_field = dict_value - assert doc._get_changed_fields() == ["dict_field"] + assert doc._get_updated_fields() == (["dict_field"], []) assert doc._delta() == ({"dict_field": dict_value}, {}) doc._changed_fields = [] list_value = ["1", 2, {"hello": "world"}] doc.list_field = list_value - assert doc._get_changed_fields() == ["list_field"] + assert doc._get_updated_fields() == (["list_field"], []) assert doc._delta() == ({"list_field": list_value}, {}) - # Test unsetting + # Test assigning empty defaults doc._changed_fields = [] doc.dict_field = {} - assert doc._get_changed_fields() == ["dict_field"] - assert doc._delta() == ({}, {"dict_field": 1}) + assert doc._get_updated_fields() == (["dict_field"], []) + assert doc._delta() == ({"dict_field": {}}, {}) doc._changed_fields = [] doc.list_field = [] - assert doc._get_changed_fields() == ["list_field"] - assert doc._delta() == ({}, {"list_field": 1}) + assert doc._get_updated_fields() == (["list_field"], []) + assert doc._delta() == ({"list_field": []}, {}) def test_delta_with_dbref_true(self): person, organization, employee = self.circular_reference_deltas_2( @@ -741,7 +1079,7 @@ def test_delta_with_dbref_true(self): ) employee.name = "test" - assert organization._get_changed_fields() == [] + assert organization._get_updated_fields() == ([], []) updates, removals = organization._delta() assert removals == {} @@ -758,7 +1096,7 @@ def test_delta_with_dbref_false(self): ) employee.name = "test" - assert organization._get_changed_fields() == [] + assert organization._get_updated_fields() == ([], []) updates, removals = organization._delta() assert removals == {} @@ -785,11 +1123,11 @@ class MyDoc(Document): subdoc = mydoc.subs["a"]["b"] subdoc.name = "bar" - assert subdoc._get_changed_fields() == ["name"] - assert mydoc._get_changed_fields() == ["subs.a.b.name"] + assert subdoc._get_updated_fields() == (["name"], []) + assert mydoc._get_updated_fields() == (["subs.a.b.name"], []) - mydoc._clear_changed_fields() - assert mydoc._get_changed_fields() == [] + mydoc._clear_updated_fields() + assert mydoc._get_updated_fields() == ([], []) def test_nested_nested_fields_db_field_set__gets_mark_as_changed_and_cleaned(self): class EmbeddedDoc(EmbeddedDocument): @@ -806,19 +1144,19 @@ class MyDoc(Document): mydoc = MyDoc.objects.first() mydoc.embed.name = "foo1" - assert mydoc.embed._get_changed_fields() == ["db_name"] - assert mydoc._get_changed_fields() == ["db_embed.db_name"] + assert mydoc.embed._get_updated_fields() == (["db_name"], []) + assert mydoc._get_updated_fields() == (["db_embed.db_name"], []) mydoc = MyDoc.objects.first() embed = EmbeddedDoc(name="foo2") embed.name = "bar" mydoc.embed = embed - assert embed._get_changed_fields() == ["db_name"] - assert mydoc._get_changed_fields() == ["db_embed"] + assert embed._get_updated_fields() == (["db_name"], []) + assert mydoc._get_updated_fields() == (["db_embed"], []) - mydoc._clear_changed_fields() - assert mydoc._get_changed_fields() == [] + mydoc._clear_updated_fields() + assert mydoc._get_updated_fields() == ([], []) def test_lower_level_mark_as_changed(self): class EmbeddedDoc(EmbeddedDocument): @@ -833,17 +1171,17 @@ class MyDoc(Document): mydoc = MyDoc.objects.first() mydoc.subs["a"] = EmbeddedDoc() - assert mydoc._get_changed_fields() == ["subs.a"] + assert mydoc._get_updated_fields() == (["subs.a"], []) subdoc = mydoc.subs["a"] subdoc.name = "bar" - assert subdoc._get_changed_fields() == ["name"] - assert mydoc._get_changed_fields() == ["subs.a"] + assert subdoc._get_updated_fields() == (["name"], []) + assert mydoc._get_updated_fields() == (["subs.a"], []) mydoc.save() - mydoc._clear_changed_fields() - assert mydoc._get_changed_fields() == [] + mydoc._clear_updated_fields() + assert mydoc._get_updated_fields() == ([], []) def test_upper_level_mark_as_changed(self): class EmbeddedDoc(EmbeddedDocument): @@ -860,15 +1198,15 @@ class MyDoc(Document): subdoc = mydoc.subs["a"] subdoc.name = "bar" - assert subdoc._get_changed_fields() == ["name"] - assert mydoc._get_changed_fields() == ["subs.a.name"] + assert subdoc._get_updated_fields() == (["name"], []) + assert mydoc._get_updated_fields() == (["subs.a.name"], []) mydoc.subs["a"] = EmbeddedDoc() - assert mydoc._get_changed_fields() == ["subs.a"] + assert mydoc._get_updated_fields() == (["subs.a"], []) mydoc.save() - mydoc._clear_changed_fields() - assert mydoc._get_changed_fields() == [] + mydoc._clear_updated_fields() + assert mydoc._get_updated_fields() == ([], []) def test_referenced_object_changed_attributes(self): """Ensures that when you save a new reference to a field, the referenced object isn't altered""" @@ -959,21 +1297,21 @@ class MyDoc(Document): MyDoc(dico={"a": {"b": 0}}).save() mydoc = MyDoc.objects.first() - assert mydoc._get_changed_fields() == [] + assert mydoc._get_updated_fields() == ([], []) mydoc.dico["a"]["b"] = 0 - assert mydoc._get_changed_fields() == [] + assert mydoc._get_updated_fields() == ([], []) mydoc.dico["a"] = {"b": 0} - assert mydoc._get_changed_fields() == [] + assert mydoc._get_updated_fields() == ([], []) mydoc.dico = {"a": {"b": 0}} - assert mydoc._get_changed_fields() == [] + assert mydoc._get_updated_fields() == ([], []) mydoc.dico["a"]["c"] = 1 - assert mydoc._get_changed_fields() == ["dico.a.c"] + assert mydoc._get_updated_fields() == (["dico.a.c"], []) mydoc.dico["a"]["b"] = 2 mydoc.dico["d"] = 3 - assert mydoc._get_changed_fields() == ["dico.a.c", "dico.a.b", "dico.d"] + assert mydoc._get_updated_fields() == (["dico.a.c", "dico.a.b", "dico.d"], []) - mydoc._clear_changed_fields() - assert mydoc._get_changed_fields() == [] + mydoc._clear_updated_fields() + assert mydoc._get_updated_fields() == ([], []) def test_delta_on_dict_empty_key_triggers_full_change(self): """more of a bug (harmless) but empty key changes aren't managed perfectly""" @@ -986,9 +1324,9 @@ class MyDoc(Document): MyDoc(dico={"a": {"b": 0}}).save() mydoc = MyDoc.objects.first() - assert mydoc._get_changed_fields() == [] + assert mydoc._get_updated_fields() == ([], []) mydoc.dico[""] = 3 - assert mydoc._get_changed_fields() == ["dico"] + assert mydoc._get_updated_fields() == (["dico"], []) mydoc.save() raw_doc = get_as_pymongo(mydoc) assert raw_doc == {"_id": mydoc.id, "dico": {"": 3, "a": {"b": 0}}} diff --git a/tests/document/test_instance.py b/tests/document/test_instance.py index 1360f825e..53ad5bf96 100644 --- a/tests/document/test_instance.py +++ b/tests/document/test_instance.py @@ -10,6 +10,7 @@ import bson import pytest from bson import SON, DBRef, ObjectId +from pymongo.collection import Collection from pymongo.errors import DuplicateKeyError, OperationFailure from mongoengine import * @@ -475,6 +476,26 @@ def test_reload(self): assert person.name == "Mr Test User" assert person.age == 21 + def test_reload__field_with_db_field_has_pending_unset__clears_pending_unset( + self, + ): + class Person(Document): + name = StringField(db_field="db_name") + + person = Person(name="stored").save() + del person.name + + assert person._get_updated_fields() == ([], ["db_name"]) + + person.reload("name") + + assert person.name == "stored" + assert person._get_updated_fields()[1] == [] + assert person._delta()[1] == {} + + person.save() + assert get_as_pymongo(person) == {"_id": person.id, "db_name": "stored"} + def test_reload_sharded(self): class Animal(Document): superphylum = StringField() @@ -571,22 +592,93 @@ class Animal(Document): def test_reload_with_changed_fields(self): """Ensures reloading will not affect changed fields""" + class VitalSigns(EmbeddedDocument): + blood_pressure = FloatField() + class User(Document): name = StringField() number = IntField() + phone = StringField() + vital_signs = EmbeddedDocumentField(VitalSigns) User.drop_collection() - user = User(name="Bob", number=1).save() + user = User( + name="Bob", + number=1, + phone="01234", + vital_signs=VitalSigns(blood_pressure=0.99), + ).save() user.name = "John" user.number = 2 + user.vital_signs.blood_pressure = 0.11 + del user.phone - assert user._get_changed_fields() == ["name", "number"] + assert user._delta() == ( + {"name": "John", "number": 2, "vital_signs.blood_pressure": 0.11}, + {"phone": 1}, + ) user.reload("number") - assert user._get_changed_fields() == ["name"] + assert user._delta() == ( + {"name": "John", "vital_signs.blood_pressure": 0.11}, + {"phone": 1}, + ) + user.reload("vital_signs") + assert user._delta() == ({"name": "John"}, {"phone": 1}) + user.reload("phone") + assert user._delta() == ({"name": "John"}, {}) user.save() + + assert get_as_pymongo(user) == { + "_id": user.id, + "name": "John", + "number": 1, + "phone": "01234", + "vital_signs": {"blood_pressure": 0.99}, + } + + del user.phone + del user.vital_signs.blood_pressure + assert user._delta() == ( + {}, + {"phone": 1, "vital_signs.blood_pressure": 1}, + ) user.reload() + assert user._delta() == ({}, {}) assert user.name == "John" + assert user.number == 1 + assert user.phone == "01234" + assert user.vital_signs.blood_pressure == 0.99 + + def test_save__reference_and_embedded_fields_are_deleted__unsets_whole_fields( + self, + ): + class VitalSigns(EmbeddedDocument): + blood_pressure = FloatField() + + class Car(Document): + brand = StringField() + + car = Car(brand="Lamborghini").save() + + class User(Document): + car = ReferenceField(Car, default=lambda: car) + vital_signs = EmbeddedDocumentField( + VitalSigns, + default=lambda: VitalSigns(blood_pressure=0.99), + ) + + user = User().save() + + del user.car + del user.vital_signs + + assert user.car == car + assert user.vital_signs.blood_pressure == 0.99 + assert user._delta() == ({}, {"car": 1, "vital_signs": 1}) + + user.save() + assert get_as_pymongo(user) == {"_id": user.id} def test_reload_referencing(self): """Ensures reloading updates weakrefs correctly.""" @@ -617,18 +709,20 @@ class Doc(Document): doc.embedded_field.list_field.append(1) doc.embedded_field.dict_field["woot"] = "woot" - changed = doc._get_changed_fields() - assert changed == [ - "list_field", - "dict_field.woot", - "embedded_field.list_field", - "embedded_field.dict_field.woot", - ] + assert doc._get_updated_fields() == ( + [ + "list_field", + "dict_field.woot", + "embedded_field.list_field", + "embedded_field.dict_field.woot", + ], + [], + ) doc.save() assert len(doc.list_field) == 4 doc = doc.reload(10) - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) assert len(doc.list_field) == 4 assert len(doc.dict_field) == 2 assert len(doc.embedded_field.list_field) == 4 @@ -638,7 +732,7 @@ class Doc(Document): doc.save() doc.dict_field["extra"] = 1 doc = doc.reload(10, "list_field") - assert doc._get_changed_fields() == ["dict_field.extra"] + assert doc._get_updated_fields() == (["dict_field.extra"], []) assert len(doc.list_field) == 5 assert len(doc.dict_field) == 3 assert len(doc.embedded_field.list_field) == 4 @@ -976,7 +1070,7 @@ def test_modify_update(self): del doc_copy.job.years assert doc.to_json() == doc_copy.to_json() - assert doc._get_changed_fields() == [] + assert doc._get_updated_fields() == ([], []) self._assert_db_equal([dict(other_doc.to_mongo()), dict(doc.to_mongo())]) @@ -1147,6 +1241,24 @@ class Person(Document): p1.reload() assert p1.name == p.parent.name + def test_save__referenced_document_has_only_unset__cascade_saves_unset(self): + class Parent(Document): + value = StringField(default="default") + + class Child(Document): + parent = ReferenceField(Parent) + + parent = Parent(value="default").save() + child = Child(parent=parent).save() + child = Child.objects.get(id=child.id) + + del child.parent.value + + assert child.parent._get_updated_fields() == ([], ["value"]) + + child.save(cascade=True) + assert get_as_pymongo(parent) == {"_id": parent.id} + def test_save_cascade_kwargs(self): class Person(Document): name = StringField() @@ -1647,7 +1759,7 @@ class User(self.Person): assert person.age == 21 assert person.active is False - def test__get_changed_fields_same_ids_reference_field_does_not_enters_infinite_loop_embedded_doc( + def test__get_updated_fields_same_ids_reference_field_does_not_enters_infinite_loop_embedded_doc( self, ): # Refers to Issue #1685 @@ -1658,10 +1770,9 @@ class ParentModel(Document): child = EmbeddedDocumentField(EmbeddedChildModel) emb = EmbeddedChildModel(id={"1": [1]}) - changed_fields = ParentModel(child=emb)._get_changed_fields() - assert changed_fields == [] + assert ParentModel(child=emb)._get_updated_fields() == ([], []) - def test__get_changed_fields_same_ids_reference_field_does_not_enters_infinite_loop_different_doc( + def test__get_updated_fields_same_ids_reference_field_does_not_enters_infinite_loop_different_doc( self, ): # Refers to Issue #1685 @@ -1680,10 +1791,10 @@ class Message(Document): message = Message(id=1, author=user).save() message.author.name = "tutu" - assert message._get_changed_fields() == [] - assert user._get_changed_fields() == ["name"] + assert message._get_updated_fields() == ([], []) + assert user._get_updated_fields() == (["name"], []) - def test__get_changed_fields_same_ids_embedded(self): + def test__get_updated_fields_same_ids_embedded(self): # Refers to Issue #1768 class User(EmbeddedDocument): id = IntField() @@ -1700,7 +1811,7 @@ class Message(Document): message = Message(id=1, author=user).save() message.author.name = "tutu" - assert message._get_changed_fields() == ["author.name"] + assert message._get_updated_fields() == (["author.name"], []) message.save() message_fetched = Message.objects.with_id(message.id) @@ -2795,6 +2906,20 @@ class PickleParent(EmbeddedDocument): assert isinstance(restored_from_raw_data.child, PickleChild) assert restored_from_raw_data.child.value == "child" + def test_pickle__legacy_state_without_tracking_fields__initializes_tracking(self): + document = PickleTest(number=1) + legacy_state = document.__getstate__() + legacy_state.pop("_changed_fields") + legacy_state.pop("_unset_fields") + legacy_state.pop("_has_change_tracking_baseline") + + restored = PickleTest.__new__(PickleTest) + restored.__setstate__(legacy_state) + + assert restored._changed_fields == [] + assert restored._unset_fields == [] + assert restored._has_change_tracking_baseline is False + def test_picklable_on_signals(self): pickle_doc = PickleSignalsTest(number=1, string="One", lists=["1", "2"]) pickle_doc.embedded = PickleEmbedded() @@ -2892,8 +3017,6 @@ class Person(Document): person = Person(name="name", age=10, job=job) - from pymongo.collection import Collection - orig_update_one = Collection.update_one try: diff --git a/tests/fields/test_dict_field.py b/tests/fields/test_dict_field.py index 98bf3d93a..245452c3d 100644 --- a/tests/fields/test_dict_field.py +++ b/tests/fields/test_dict_field.py @@ -21,6 +21,45 @@ class BlogPost(Document): post = BlogPost(info=info).save() assert get_as_pymongo(post) == {"_id": post.id, "info": info} + def test_save__dict_key_is_assigned_none__stores_null(self): + """Regression test for issue #2051.""" + + class Test(Document): + args = DictField() + + test = Test(args={"hello": None}).save() + + test.args["count"] = None + + assert test._delta() == ({"args.count": None}, {}) + + test.save() + assert get_as_pymongo(test) == { + "_id": test.id, + "args": {"hello": None, "count": None}, + } + + def test_save__embedded_dict_key_is_assigned_none__stores_null(self): + """Regression test for issue #1378.""" + + class Item(EmbeddedDocument): + field = DictField() + + class Doc(Document): + items = ListField(EmbeddedDocumentField(Item)) + + doc = Doc(items=[Item()]).save() + + doc.items[0].field["key"] = None + + assert doc._delta() == ({"items.0.field.key": None}, {}) + + doc.save() + assert get_as_pymongo(doc) == { + "_id": doc.id, + "items": [{"field": {"key": None}}], + } + def test_validate_invalid_type(self): class BlogPost(Document): info = DictField() diff --git a/tests/fields/test_generic_reference_field.py b/tests/fields/test_generic_reference_field.py index 6609fb32e..e0fa76574 100644 --- a/tests/fields/test_generic_reference_field.py +++ b/tests/fields/test_generic_reference_field.py @@ -348,10 +348,10 @@ class Doc2(Document): doc2 = Doc2(ref=doc1, refs=[doc11]).save() doc2.ref.name = "garbage2" - assert doc2._get_changed_fields() == [] + assert doc2._get_updated_fields() == ([], []) doc2.refs[0].name = "garbage3" - assert doc2._get_changed_fields() == [] + assert doc2._get_updated_fields() == ([], []) assert doc2._delta() == ({}, {}) def test_generic_reference_field(self): diff --git a/tests/queryset/test_field_list.py b/tests/queryset/test_field_list.py index 25a7c7619..b0bc9ecc1 100644 --- a/tests/queryset/test_field_list.py +++ b/tests/queryset/test_field_list.py @@ -4,6 +4,7 @@ from mongoengine import * from mongoengine.queryset import QueryFieldList +from tests.utils import MongoDBTestCase, get_as_pymongo class TestQueryFieldList: @@ -66,6 +67,72 @@ def test_using_a_slice(self): assert q.as_dict() == {"a": {"$slice": 5}} +class TestListField(MongoDBTestCase): + def test_save__list_field_is_empty__stores_empty_list_and_del_unsets(self): + class BlogPost(Document): + authors = ListField(default=[]) + + BlogPost.drop_collection() + + blog = BlogPost().save() + assert get_as_pymongo(blog) == {"_id": blog.id, "authors": []} + + blog.authors = [] + blog.save() + assert get_as_pymongo(blog) == {"_id": blog.id, "authors": []} + + blog.authors = [1] + blog.save() + assert get_as_pymongo(blog) == {"_id": blog.id, "authors": [1]} + + del blog.authors + blog.save() + assert get_as_pymongo(blog) == {"_id": blog.id} + + blog = BlogPost(authors=[]).save() + assert get_as_pymongo(blog) == {"_id": blog.id, "authors": []} + + blog = BlogPost(authors=None).save() + assert get_as_pymongo(blog) == {"_id": blog.id, "authors": []} + + def test_only__list_field_is_missing__returns_default(self): + # Ensure no regression of #938 + class Doc(Document): + values = ListField(IntField()) + + Doc.drop_collection() + + doc = Doc(values=[]).save() + del doc.values + doc.save() + + assert get_as_pymongo(doc) == {"_id": doc.id} + assert Doc.objects(id=doc.id).only("values").get().values == [] + + def test_item_frequencies__list_is_empty_or_missing__distinguishes_values(self): + class Doc(Document): + fruit = ListField(StringField()) + + Doc.drop_collection() + + doc = Doc(fruit=["a", "a", "b"]).save() + other_doc = Doc(fruit=["b", "c"]).save() + + assert Doc.objects.item_frequencies("fruit") == {"a": 2, "b": 2, "c": 1} + + other_doc.delete() + assert Doc.objects.item_frequencies("fruit") == {"a": 2, "b": 1} + + doc.fruit = [] + doc.save() + assert Doc.objects.item_frequencies("fruit") == {} + + del doc.fruit + doc.save() + assert get_as_pymongo(doc) == {"_id": doc.id} + assert Doc.objects.item_frequencies("fruit") == {None: 1} + + class TestOnlyExcludeAll(unittest.TestCase): def setUp(self): connect(db="mongoenginetest") diff --git a/tests/queryset/test_queryset.py b/tests/queryset/test_queryset.py index fae2c0ff8..0e6604e45 100644 --- a/tests/queryset/test_queryset.py +++ b/tests/queryset/test_queryset.py @@ -905,13 +905,13 @@ class TestOrganization(Document): o.owner = p p.name = "p2" - assert o._get_changed_fields() == ["owner"] - assert p._get_changed_fields() == ["name"] + assert o._get_updated_fields() == (["owner"], []) + assert p._get_updated_fields() == (["name"], []) o.save() - assert o._get_changed_fields() == [] - assert p._get_changed_fields() == ["name"] # Fails; it's empty + assert o._get_updated_fields() == ([], []) + assert p._get_updated_fields() == (["name"], []) # Fails; it's empty # This will do NOTHING at all, even though we changed the name p.save() @@ -1185,7 +1185,7 @@ class Comment(Document): with pytest.raises(NotUniqueError): Comment.objects.insert(com1) - def test_get_changed_fields_query_count(self): + def test_get_updated_fields_query_count(self): """Make sure we don't perform unnecessary db operations when none of document's fields were updated. """ @@ -1223,7 +1223,7 @@ class Project(Document): # Checking changed fields of a newly fetched document should not # result in a query. - org._get_changed_fields() + org._get_updated_fields() assert q == 1 # Saving a doc without changing any of its fields should not result diff --git a/tests/test_datastructures.py b/tests/test_datastructures.py index eb1417d11..209926103 100644 --- a/tests/test_datastructures.py +++ b/tests/test_datastructures.py @@ -73,10 +73,11 @@ def test_clear_calls_mark_as_changed(self): assert base_dict._instance._changed_fields == ["my_name"] assert base_dict == {} - def test___delitem___calls_mark_as_changed(self): + def test___delitem___calls_mark_as_unset(self): base_dict = self._get_basedict({"k": "v"}) del base_dict["k"] - assert base_dict._instance._changed_fields == ["my_name.k"] + assert base_dict._instance._changed_fields == [] + assert base_dict._instance._unset_fields == ["my_name.k"] assert base_dict == {} def test___getitem____KeyError(self): @@ -148,14 +149,21 @@ def test___setattr____not_tracked_by_changes(self): base_dict.a_new_attr = "test" assert base_dict._instance._changed_fields == [] - def test___delattr____tracked_by_changes(self): - # This is probably a bug as __setattr__ is not tracked - # This is even bad because it could be that there is an attribute - # with the same name as a key + def test___delattr____does_not_track_changes(self): base_dict = self._get_basedict({}) base_dict.a_new_attr = "test" del base_dict.a_new_attr - assert base_dict._instance._changed_fields == ["my_name.a_new_attr"] + assert base_dict._instance._changed_fields == [] + assert base_dict._instance._unset_fields == [] + + def test___delattr____missing_attribute_does_not_track_changes(self): + base_dict = self._get_basedict({}) + + with pytest.raises(AttributeError): + del base_dict.missing + + assert base_dict._instance._changed_fields == [] + assert base_dict._instance._unset_fields == [] class TestBaseList: diff --git a/tests/test_dereference.py b/tests/test_dereference.py index 224538312..d054a4fd4 100644 --- a/tests/test_dereference.py +++ b/tests/test_dereference.py @@ -416,7 +416,7 @@ def __repr__(self): daughter.relations.append(mother) daughter.relations.append(daughter) - assert daughter._get_changed_fields() == ["relations"] + assert daughter._get_updated_fields() == (["relations"], []) daughter.save() assert "[, ]" == "%s" % Person.objects()