diff --git a/docs/changelog.rst b/docs/changelog.rst index bfaeab77f..53a27f3f5 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -7,7 +7,6 @@ Changelog Development =========== - (Fill this out as you fix issues and develop your features). -- Fix partial ``Document.reload()`` leaving custom ``db_field`` values marked as changed. Changes in 1.0.0 ================ @@ -42,6 +41,8 @@ Changes in 1.0.0 - Bug Fix - Fix querying GenericReferenceField with __in operator #2886 - Bug Fix - Fix Document.compare_indexes() not working correctly for text indexes on multiple fields #2612 - BREAKING CHANGE: wrap _document_registry (normally not used by end users) with _DocumentRegistry which acts as a singleton to access the registry +- Fix partial ``Document.reload()`` leaving custom ``db_field`` values marked as changed. +- Fix ``ListField`` and ``DictField`` failing to retain ``None`` values for non-required inner fields. #2290 #913 #1904 - Log a warning in case users creates multiple Document classes with the same name as it can lead to unexpected behavior #1778 - Fix use of $search, $searchMeta, or $vectorSearch in aggregate #2878 - BugFix - Fix use of $geoNear or $collStats in aggregate #2493 diff --git a/mongoengine/base/fields.py b/mongoengine/base/fields.py index 6da31ff85..b1bd648e1 100644 --- a/mongoengine/base/fields.py +++ b/mongoengine/base/fields.py @@ -266,6 +266,11 @@ def _validate_choices(self, value): self.error("Value must be one of %s" % str(choice_list)) def _validate(self, value, **kwargs): + if value is None: + if self.required: + self.error("Field is required") + return + # Check the Choices Constraint if self.choices: self._validate_choices(value) @@ -427,7 +432,8 @@ def to_python(self, value): if self.field: self.field.set_auto_dereferencing(self._auto_dereference) value_dict = { - key: self.field.to_python(item) for key, item in value.items() + key: None if item is None else self.field.to_python(item) + for key, item in value.items() } else: Document = _import_class("Document") @@ -482,7 +488,11 @@ def to_mongo(self, value, use_db_field=True, fields=None): if self.field: value_dict = { - key: self.field._to_mongo_safe_call(item, use_db_field, fields) + key: ( + None + if item is None + else self.field._to_mongo_safe_call(item, use_db_field, fields) + ) for key, item in value.items() } else: diff --git a/mongoengine/fields.py b/mongoengine/fields.py index 37d5070a0..06e8aa71a 100644 --- a/mongoengine/fields.py +++ b/mongoengine/fields.py @@ -975,7 +975,10 @@ def prepare_query_value(self, op, value): op in ("set", "unset", "gt", "gte", "lt", "lte", "ne", None) and eligible_iter ): - return [self.field.prepare_query_value(op, v) for v in value] + return [ + None if v is None else self.field.prepare_query_value(op, v) + for v in value + ] return self.field.prepare_query_value(op, value) @@ -1025,6 +1028,11 @@ def to_mongo(self, value, use_db_field=True, fields=None): ) return sorted(value, reverse=self._order_reverse) + def validate(self, value): + super().validate(value) + if any(item is None for item in value): + self.error("SortedListField does not support None values") + def key_not_string(d): """Helper function to recursively determine if any key in a @@ -1091,7 +1099,8 @@ def prepare_query_value(self, op, value): ): # Used for instance when using DictField(ListField(IntField())) if op in ("set", "unset") and isinstance(value, dict): return { - k: self.field.prepare_query_value(op, v) for k, v in value.items() + k: None if v is None else self.field.prepare_query_value(op, v) + for k, v in value.items() } return self.field.prepare_query_value(op, value) diff --git a/mongoengine/queryset/transform.py b/mongoengine/queryset/transform.py index 7d2c6e495..007ab909a 100644 --- a/mongoengine/queryset/transform.py +++ b/mongoengine/queryset/transform.py @@ -352,15 +352,24 @@ def update(_doc_cls=None, **update): else: value = field.prepare_query_value(op, value) elif op == "push" and isinstance(value, (list, tuple, set)): - value = [field.prepare_query_value(op, v) for v in value] + value = [ + None if v is None else field.prepare_query_value(op, v) + for v in value + ] elif op in (None, "set", "push"): if field.required or value is not None: value = field.prepare_query_value(op, value) elif op in ("pushAll", "pullAll"): - value = [field.prepare_query_value(op, v) for v in value] + value = [ + None if v is None else field.prepare_query_value(op, v) + for v in value + ] elif op in ("addToSet", "setOnInsert"): if isinstance(value, (list, tuple, set)): - value = [field.prepare_query_value(op, v) for v in value] + value = [ + None if v is None else field.prepare_query_value(op, v) + for v in value + ] elif field.required or value is not None: value = field.prepare_query_value(op, value) elif op == "unset": diff --git a/tests/document/test_instance.py b/tests/document/test_instance.py index 9b4d3b672..0785f8d82 100644 --- a/tests/document/test_instance.py +++ b/tests/document/test_instance.py @@ -2394,7 +2394,7 @@ class Word(Document): "stem": [1, 2, 3], "forms": 1, "count": "one", - "occurs": {"hello": None}, + "occurs": {"hello": "not_a_dict"}, } ) diff --git a/tests/fields/test_dict_field.py b/tests/fields/test_dict_field.py index 245452c3d..16819c28c 100644 --- a/tests/fields/test_dict_field.py +++ b/tests/fields/test_dict_field.py @@ -39,6 +39,22 @@ class Test(Document): "args": {"hello": None, "count": None}, } + def test_save__typed_dict_value_is_none__retains_none(self): + class Test(Document): + values = DictField(IntField()) + + test = Test(values={"missing": None}).save() + test.reload() + + assert test.values == {"missing": None} + + def test_validate__typed_dict_required_value_is_none__raises_required_error(self): + class Test(Document): + values = DictField(IntField(required=True)) + + with pytest.raises(ValidationError, match="Field is required"): + Test(values={"missing": None}).validate() + def test_save__embedded_dict_key_is_assigned_none__stores_null(self): """Regression test for issue #1378.""" diff --git a/tests/fields/test_fields.py b/tests/fields/test_fields.py index e6f0f7590..429c50ad5 100644 --- a/tests/fields/test_fields.py +++ b/tests/fields/test_fields.py @@ -564,6 +564,53 @@ class BlogPost(Document): post.generic_as_lazy = [user] post.validate() + def test_list_field__optional_string_items_are_none__retains_none(self): + class BlogPost(Document): + tags = ListField(StringField()) + + post = BlogPost(tags=[None, "hello", None]).save() + post.reload() + + assert post.tags == [None, "hello", None] + + def test_list_field__nested_optional_string_item_is_none__retains_none(self): + class BlogPost(Document): + tags = ListField(ListField(StringField())) + + post = BlogPost(tags=[[None]]).save() + post.reload() + + assert post.tags == [[None]] + + def test_list_field__optional_embedded_items_are_none__retains_none(self): + class Comment(EmbeddedDocument): + content = StringField() + + class BlogPost(Document): + comments = ListField(EmbeddedDocumentField(Comment)) + + post = BlogPost(comments=[None, Comment(content="hello")]).save() + post.reload() + + assert post.comments[0] is None + assert post.comments[1].content == "hello" + + def test_list_field__required_item_is_none__raises_required_error(self): + class BlogPost(Document): + tags = ListField(StringField(required=True)) + + with pytest.raises(ValidationError, match="Field is required"): + BlogPost(tags=[None]).validate() + + def test_sorted_list_field__item_is_none__raises_validation_error(self): + class BlogPost(Document): + tags = SortedListField(StringField()) + + with pytest.raises( + ValidationError, match="SortedListField does not support None values" + ): + BlogPost(tags=[None, "hello"]).validate() + def test_sorted_list_sorting(self): """Ensure that a sorted list field properly sorts values.""" diff --git a/tests/queryset/test_transform.py b/tests/queryset/test_transform.py index 75e024b95..eab3b2b24 100644 --- a/tests/queryset/test_transform.py +++ b/tests/queryset/test_transform.py @@ -110,6 +110,19 @@ class BlogPost(Document): update = transform.update(BlogPost, push_all__tags=["mongo", "db"]) assert update == {"$push": {"tags": {"$each": ["mongo", "db"]}}} + def test_transform_update__list_items_are_none__retains_none(self): + class BlogPost(Document): + tags = ListField(StringField()) + + for operator in ("set", "push"): + update = transform.update( + BlogPost, **{f"{operator}__tags": [None, "hello", None]} + ) + + assert update == { + f"${operator}": {"tags": [None, "hello", None]}, + } + def test_transform_update_inc_dec_ignores_min_max(self): """inc/dec pass a delta; min_value/max_value apply to stored values (#2339)."""