diff --git a/mongodol/base.py b/mongodol/base.py index 8a8523e..755463a 100644 --- a/mongodol/base.py +++ b/mongodol/base.py @@ -1,11 +1,14 @@ """Base mongoDB data object layers""" +import re from functools import wraps, cached_property from typing import Optional, Union from collections.abc import Mapping from collections import ChainMap from dol.base import Store +import bson +from bson.regex import Regex as BsonRegex from pymongo import MongoClient from dol import KvReader, Collection as DolCollection @@ -25,6 +28,58 @@ ) +_QUERY_VALUE_TYPES = (re.Pattern, BsonRegex) + + +def operator_field_names(obj) -> list: + """The parts of ``obj`` that make it act as a query rather than an exact match. + + That is: ``$``-prefixed field names and regular-expression values, found at any + depth (in mappings and lists). Regexes are reported by their ``repr``. + + >>> operator_field_names({'a': 1, 'b': {'c': [{'$gt': 2}]}}) + ['$gt'] + >>> operator_field_names({'a': re.compile('x')}) + ["re.compile('x')"] + >>> operator_field_names({'a': 1}) + [] + """ + found = [] + if isinstance(obj, _QUERY_VALUE_TYPES): + found.append(repr(obj)) + elif isinstance(obj, Mapping): + for field, value in obj.items(): + if isinstance(field, str) and field.startswith("$"): + found.append(field) + found.extend(operator_field_names(value)) + elif isinstance(obj, (list, tuple)): + for value in obj: + found.extend(operator_field_names(value)) + return found + + +def _is_operator_expression(value) -> bool: + return isinstance(value, Mapping) and any( + isinstance(f, str) and f.startswith("$") for f in value + ) + + +def _same_bson_value(a, b) -> bool: + """Equality as MongoDB sees it (type- and field-order-sensitive), not Python's.""" + return bson.encode({"v": a}) == bson.encode({"v": b}) + + +def _satisfies_scope_value(value, scope_value) -> Optional[bool]: + """Whether ``value`` satisfies ``scope_value``; ``None`` if that can't be decided here.""" + if not _is_operator_expression(scope_value): + return _same_bson_value(value, scope_value) + if set(scope_value) == {"$eq"}: + return _same_bson_value(value, scope_value["$eq"]) + if set(scope_value) == {"$in"}: + return any(_same_bson_value(value, x) for x in scope_value["$in"]) + return None # other operators: not checked (use on_write_filter for such scopes) + + # TODO: mgc type annotation # See https://stackoverflow.com/questions/66464191/referencing-a-python-class-within-its-definition-but-outside-a-method class MongoCollectionCollection(DolCollection): @@ -402,8 +457,21 @@ class MongoCollectionPersister(MongoCollectionReader): {'first': 'Guido', 'last': 'van Rossum'} --> {'yob': 1956, 'proj': 'python', 'bdfl': False} {'first': 'Vitalik', 'last': 'Buterin'} --> {'yob': 1994, 'proj': 'ethereum', 'bdfl': True} + Writes stay inside the store's scope: a key or value that contradicts a field + of the write filter (``on_write_filter``, else ``filter``) raises + ``ValueError``. Fields scoped with operators other than ``$eq``/``$in`` (such + as ``$ne``, ``$gt``) are NOT checked and such writes are let through: give + those stores an ``on_write_filter`` with plain values. Keys used to replace or + delete docs may not contain ``$``-operators or regexes (pass + ``allow_operators_in_write_keys=True``, or set it as a class attribute, to allow + them), and those queries are confined by ``filter`` and ``on_write_filter``. + Reads (``s[k]``, ``k in s``) still accept query keys, always within ``filter``. + """ + #: Whether keys given to write/delete operations may contain ``$``-operators. + allow_operators_in_write_keys = False + def __init__( self, mgc: PyMongoCollectionSpec | KvReader = None, @@ -411,6 +479,8 @@ def __init__( on_write_filter: dict | None = None, iter_projection: ProjectionSpec = (ID,), getitem_projection: ProjectionSpec = None, + *, + allow_operators_in_write_keys: Optional[bool] = None, **mgc_find_kwargs, ): super().__init__( @@ -421,13 +491,15 @@ def __init__( **mgc_find_kwargs, ) self._on_write_filter = on_write_filter + if allow_operators_in_write_keys is not None: + self.allow_operators_in_write_keys = allow_operators_in_write_keys def __setitem__(self, k, v): assert isinstance(k, Mapping) and isinstance(v, Mapping), ( f"k (key) and v (value) must both be mappings (often dictionaries). Were:\n\tk={k}\n\tv={v}" ) return self.mgc.replace_one( - filter=self._merge_with_filt(k), + filter=self._write_filter_for_key(k), replacement=self._build_doc(k, v), upsert=True, ) @@ -437,10 +509,30 @@ def __delitem__(self, k): f"k (key) must be a mapping (most often a dictionary). Were:\n\tk={k}" ) if len(k) > 0: - return self.mgc.delete_one(self._merge_with_filt(k)) + return self.mgc.delete_one(self._write_filter_for_key(k)) else: raise KeyError(f"You can't remove that key: {k}") + def _write_filter_for_key(self, k: Mapping) -> dict: + """The query selecting the doc(s) that a write/delete of key ``k`` targets. + + Refuses query-shaped keys (``$``-operators, regexes) unless + ``allow_operators_in_write_keys``, since such a key would select arbitrary + docs of the scope. The query is confined by ``filter`` and, when set, + ``on_write_filter``. + """ + if not self.allow_operators_in_write_keys: + operators = operator_field_names(k) + if operators: + raise ValueError( + f"Keys used to write or delete may not contain query operators " + f"or patterns ({', '.join(sorted(set(operators)))}). Key was: {k}" + ) + on_write_filter = getattr(self, "_on_write_filter", None) + if on_write_filter: + return {"$and": [self.filter, on_write_filter, k]} + return self._merge_with_filt(k) + def append(self, v): """Insert a single doc ``v``, merged with ``on_write_filter`` if set, else this store's filter.""" assert isinstance(v, Mapping), ( @@ -458,13 +550,28 @@ def extend(self, values): def _build_doc(self, *args): def merge_doc_elements_with_filter(): - d = self._on_write_filter or self.filter + scope = self._on_write_filter or self.filter + dotted = [f for f in scope if isinstance(f, str) and "." in f] + if dotted: + raise ValueError( + f"Can't write through a scope with dotted fields {dotted}: " + "give an on_write_filter with nested documents instead." + ) + d = dict(scope) for v in args: if v is None: v = {} assert isinstance(v, Mapping), ( f" v (value) must be a mapping (often a dictionary). Were:\n\tv={v}" ) + for field, value in v.items(): + if field in scope and _satisfies_scope_value( + value, scope[field] + ) is False: + raise ValueError( + f"Field {field!r} is {value!r}, which contradicts this " + f"store's write scope ({field!r}: {scope[field]!r})." + ) d = dict(d, **v) return d diff --git a/mongodol/stores.py b/mongodol/stores.py index 88ec1be..eb608c7 100644 --- a/mongodol/stores.py +++ b/mongodol/stores.py @@ -204,11 +204,14 @@ def __setitem__(self, k, v): ), ( f"v (value) must be mappings (often dictionaries) or a collection of mappings. Were:\n\tk={k}\n\tv={v}" ) - self.mgc.delete_many(self._merge_with_filt(k)) # A Mapping is itself a Collection, so it must be tested for first, or a single # doc would be "iterated" into its field names. docs = [v] if isinstance(v, Mapping) else list(v) - return self.mgc.insert_many([self._build_doc(k, doc) for doc in docs]) + # Validate everything before deleting anything, so a refused write loses no data. + delete_filter = self._write_filter_for_key(k) + new_docs = [self._build_doc(k, doc) for doc in docs] + self.mgc.delete_many(delete_filter) + return self.mgc.insert_many(new_docs) class MongoStore(Store): diff --git a/mongodol/tests/write_scope_test.py b/mongodol/tests/write_scope_test.py new file mode 100644 index 0000000..3a9fc9d --- /dev/null +++ b/mongodol/tests/write_scope_test.py @@ -0,0 +1,166 @@ +"""Writes stay within a store's scope filter; write/delete keys can't carry operators. + +These tests use a recording stand-in for the pymongo collection, so they need no +server (but live under ``tests/``, which conftest deselects without one). +""" + +import re + +import pytest +from bson.regex import Regex +from pymongo import MongoClient + +from mongodol.base import MongoCollectionPersister, operator_field_names +from mongodol.stores import MongoCollectionMultipleDocsPersister + + +class _RecordingCollection: + """Records every call made to it; no server involved.""" + + def __init__(self): + self.calls = [] + + def __getattr__(self, name): + def method(*args, **kwargs): + self.calls.append((name, args, kwargs)) + + return method + + +def _store(cls=MongoCollectionPersister, **kwargs): + # A real (never connected) collection to construct with, then swap in the recorder. + unconnected = MongoClient("mongodb://localhost:1", connect=False)["db"]["c"] + s = cls(unconnected, **kwargs) + mgc = _RecordingCollection() + # dol-wrapped classes (e.g. via wrap_kvs) hold the persister in ``.store`` + getattr(s, "store", s).mgc = mgc + return s, mgc + + +def test_operator_field_names(): + assert operator_field_names({"a": {"$ne": 1}}) == ["$ne"] + assert operator_field_names({"$or": [{"a": 1}, {"b": {"$gt": 2}}]}) == [ + "$or", + "$gt", + ] + assert operator_field_names({"a": [1, {"b": 2}]}) == [] + + +@pytest.mark.parametrize( + "bad_key", + [ + {"_id": {"$ne": None}}, + {"$or": [{"_id": 1}, {"_id": 2}]}, + {"_id": re.compile(".*")}, + {"name": Regex("")}, + ], +) +def test_write_and_delete_refuse_operator_keys(bad_key): + s, mgc = _store(filter={"tenant": "a"}) + with pytest.raises(ValueError, match="query operators"): + s[bad_key] = {"x": 1} + with pytest.raises(ValueError, match="query operators"): + del s[bad_key] + assert mgc.calls == [] + + +@pytest.mark.parametrize("bad_key", [{"_id": {"$exists": True}}, {"g": re.compile("")}]) +def test_multiple_docs_setitem_refuses_operator_keys_before_delete_many(bad_key): + s, mgc = _store(MongoCollectionMultipleDocsPersister, filter={"tenant": "a"}) + with pytest.raises(ValueError, match="query operators"): + s[bad_key] = [{"x": 1}] + assert mgc.calls == [] + + +def test_multiple_docs_setitem_validates_values_before_delete_many(): + s, mgc = _store(MongoCollectionMultipleDocsPersister, filter={"tenant": "a"}) + with pytest.raises(ValueError, match="contradicts"): + s[{"g": 1}] = [{"x": 1}, {"x": 2, "tenant": "b"}] + assert mgc.calls == [] + s[{"g": 1}] = [{"x": 1}] + assert [name for name, *_ in mgc.calls] == ["delete_many", "insert_many"] + + +def test_class_attribute_opt_in(): + class Permissive(MongoCollectionPersister): + allow_operators_in_write_keys = True + + s, mgc = _store(Permissive) + del s[{"_id": {"$ne": None}}] + assert [name for name, *_ in mgc.calls] == ["delete_one"] + + +def test_operator_keys_allowed_with_explicit_opt_in(): + s, mgc = _store(filter={"tenant": "a"}, allow_operators_in_write_keys=True) + del s[{"_id": {"$ne": None}}] + assert [name for name, *_ in mgc.calls] == ["delete_one"] + + +@pytest.mark.parametrize( + "k, v", + [({"_id": 1, "tenant": "b"}, {"x": 1}), ({"_id": 1}, {"x": 1, "tenant": "b"})], +) +def test_contradicting_scope_field_is_refused(k, v): + s, mgc = _store(filter={"tenant": "a"}) + with pytest.raises(ValueError, match="contradicts"): + s[k] = v + with pytest.raises(ValueError, match="contradicts"): + s.append(dict(k, **v)) + assert mgc.calls == [] + + +def test_consistent_writes_unchanged(): + s, mgc = _store(filter={"tenant": "a"}) + s[{"_id": 1, "tenant": "a"}] = {"x": 1} + s[{"_id": 2}] = {"x": 2} + (name, _, kwargs), (_, _, kwargs2) = mgc.calls + assert name == "replace_one" + assert kwargs["replacement"] == {"tenant": "a", "_id": 1, "x": 1} + assert kwargs["filter"] == {"$and": [{"tenant": "a"}, {"_id": 1, "tenant": "a"}]} + assert kwargs2["replacement"] == {"tenant": "a", "_id": 2, "x": 2} + + +def test_on_write_filter_is_the_write_scope(): + s, mgc = _store(filter={"tenant": {"$in": ["a", "b"]}}, on_write_filter={"tenant": "a"}) + s.append({"x": 1}) + assert mgc.calls[0][1][0] == {"tenant": "a", "x": 1} + with pytest.raises(ValueError, match="contradicts"): + s.append({"x": 1, "tenant": "b"}) + + +@pytest.mark.parametrize( + "scope, v", + [ + ({"active": True}, {"active": 1}), # equal in Python, not in MongoDB + ({"m": {"a": 1, "b": 2}}, {"m": {"b": 2, "a": 1}}), # field order matters + ], +) +def test_contradiction_uses_mongodb_equality(scope, v): + s, mgc = _store(filter=scope) + with pytest.raises(ValueError, match="contradicts"): + s.append(v) + assert mgc.calls == [] + + +def test_operator_scope_writes(): + s, mgc = _store(filter={"tenant": {"$in": ["a", "b"]}}) + s[{"_id": 1, "tenant": "a"}] = {"x": 1} # inside the scope: allowed, as before + assert mgc.calls[0][2]["replacement"]["tenant"] == "a" + with pytest.raises(ValueError, match="contradicts"): + s[{"_id": 2, "tenant": "c"}] = {"x": 1} + + +def test_dotted_scope_fields_refuse_writes(): + s, mgc = _store(filter={"org.id": "a"}) + with pytest.raises(ValueError, match="dotted"): + s.append({"x": 1}) + assert mgc.calls == [] + + +def test_on_write_filter_confines_replace_and_delete(): + s, mgc = _store(on_write_filter={"tenant": "a"}) + del s[{"_id": 1}] + s[{"_id": 2}] = {"x": 1} + (_, delete_args, _), (_, _, replace_kwargs) = mgc.calls + assert delete_args[0] == {"$and": [{}, {"tenant": "a"}, {"_id": 1}]} + assert replace_kwargs["filter"] == {"$and": [{}, {"tenant": "a"}, {"_id": 2}]} diff --git a/mongodol/tracking_methods.py b/mongodol/tracking_methods.py index 36263f9..c152cb0 100644 --- a/mongodol/tracking_methods.py +++ b/mongodol/tracking_methods.py @@ -270,12 +270,12 @@ def get_op_request(func, *args, **kwargs): if func_name == "__setitem__": k = _kwargs.get("k", {}) return pymongo.ReplaceOne( - filter=self._merge_with_filt(k), + filter=self._write_filter_for_key(k), replacement=self._build_doc(k, v), upsert=True, ) elif func_name == "__delitem__": - return pymongo.DeleteOne(filter=self._merge_with_filt(k)) + return pymongo.DeleteOne(filter=self._write_filter_for_key(k)) elif func_name == "append": return pymongo.InsertOne(document=self._build_doc(v)) elif func_name == "extend":