diff --git a/README.md b/README.md index 95340e0..90249d4 100644 --- a/README.md +++ b/README.md @@ -54,6 +54,14 @@ m = DotMap() m.people.steve.age = 31 ``` +You can provide a factory for default values of missing keys when you do not want automatic hierarchy creation for absent keys. Like `collections.defaultdict`, the factory is called on each miss and the result is stored under the key + +```python +m = DotMap({'city': 'abc', 'CountryCode': 101}, _default_factory=str) +print(m.zipCode) +# '' +``` + And key initialization ```python diff --git a/dotmap/__init__.py b/dotmap/__init__.py index 25da0b0..ddcb84d 100755 --- a/dotmap/__init__.py +++ b/dotmap/__init__.py @@ -19,6 +19,11 @@ class DotMap(MutableMapping, OrderedDict): def __init__(self, *args, **kwargs): self._map = OrderedDict() self._dynamic = kwargs.pop('_dynamic', True) + self._default_factory = kwargs.pop('_default_factory', None) + if self._default_factory is not None and not callable(self._default_factory): + raise TypeError('_default_factory must be callable') + if self._default_factory is not None and not self._dynamic: + raise ValueError('cannot provide _default_factory when _dynamic is False') self._prevent_method_masking = kwargs.pop('_prevent_method_masking', False) _key_convert_hook = kwargs.pop('_key_convert_hook', None) @@ -46,7 +51,14 @@ def __init__(self, *args, **kwargs): v = trackedIDs[idv] else: trackedIDs[idv] = v - v = self.__class__(v, _dynamic=self._dynamic, _prevent_method_masking = self._prevent_method_masking, _key_convert_hook =_key_convert_hook, _trackedIDs = trackedIDs) + child_kwargs = { + '_dynamic': self._dynamic, + '_default_factory': self._default_factory, + '_prevent_method_masking': self._prevent_method_masking, + '_key_convert_hook': _key_convert_hook, + '_trackedIDs': trackedIDs + } + v = self.__class__(v, **child_kwargs) if type(v) is list: l = [] for i in v: @@ -57,7 +69,13 @@ def __init__(self, *args, **kwargs): n = trackedIDs[idi] else: trackedIDs[idi] = i - n = self.__class__(i, _dynamic=self._dynamic, _key_convert_hook =_key_convert_hook, _prevent_method_masking = self._prevent_method_masking) + child_kwargs = { + '_dynamic': self._dynamic, + '_default_factory': self._default_factory, + '_key_convert_hook': _key_convert_hook, + '_prevent_method_masking': self._prevent_method_masking + } + n = self.__class__(i, **child_kwargs) l.append(n) v = l self._map[k] = v @@ -90,13 +108,21 @@ def next(self): def __setitem__(self, k, v): self._map[k] = v def __getitem__(self, k): - if k not in self._map and self._dynamic and k != '_ipython_canary_method_should_not_exist_': - # automatically extend to new DotMap - self[k] = self.__class__() + if k not in self._map: + if self._dynamic and k != '_ipython_canary_method_should_not_exist_': + if self._default_factory is not None: + self[k] = self._default_factory() + else: + # automatically extend to new DotMap + self[k] = self.__class__() return self._map[k] def __setattr__(self, k, v): - if k in {'_map','_dynamic', '_ipython_canary_method_should_not_exist_', '_prevent_method_masking'}: + if k in { + '_map', '_dynamic', '_default_factory', + '_ipython_canary_method_should_not_exist_', + '_prevent_method_masking' + }: super(DotMap, self).__setattr__(k,v) elif self._prevent_method_masking and k in reserved_keys: raise KeyError('"{}" is reserved'.format(k)) @@ -107,7 +133,10 @@ def __getattr__(self, k): if k.startswith('__') and k.endswith('__'): raise AttributeError(k) - if k in {'_map','_dynamic','_ipython_canary_method_should_not_exist_'}: + if k in { + '_map', '_dynamic', '_default_factory', + '_ipython_canary_method_should_not_exist_' + }: return super(DotMap, self).__getattr__(k) try: @@ -243,7 +272,7 @@ def __len__(self): def clear(self): self._map.clear() def copy(self): - return self.__class__(self) + return self.__class__(self, _default_factory=self._default_factory) def __copy__(self): return self.copy() def __deepcopy__(self, memo=None): @@ -280,7 +309,10 @@ def fromkeys(cls, seq, value=None): d._map = OrderedDict.fromkeys(seq, value) return d def __getstate__(self): return self.__dict__ - def __setstate__(self, d): self.__dict__.update(d) + def __setstate__(self, d): + self.__dict__.update(d) + if '_default_factory' not in self.__dict__: + self._default_factory = None # bannerStr def _getListStr(self,items): out = '[' diff --git a/dotmap/test.py b/dotmap/test.py index 8b10f90..b1bf8d8 100644 --- a/dotmap/test.py +++ b/dotmap/test.py @@ -1,4 +1,5 @@ import unittest +import copy from dotmap import DotMap @@ -192,6 +193,70 @@ def assignNonDynamicKeyWithInit(): self.assertRaises(KeyError, assignNonDynamicKeyWithInit) +class TestDefault(unittest.TestCase): + def test_missing_attribute_returns_default(self): + address = {'city': 'abc', 'country': 'XY', 'CountryCode': 101} + m = DotMap(address, _default_factory=str) + + self.assertEqual(m.city, 'abc') + self.assertEqual(m.CountryCode, 101) + self.assertNotIn('zipCode', m) + self.assertEqual(m.zipCode, '') + + def test_nested_maps_inherit_default(self): + m = DotMap({'address': {'city': 'abc'}}, _default_factory=str) + + self.assertEqual(m.address.city, 'abc') + self.assertEqual(m.address.zipCode, '') + + def test_default_with_dynamic_false_raises(self): + with self.assertRaises(ValueError): + DotMap(_default_factory=str, _dynamic=False) + + def test_default_factory_must_be_callable(self): + with self.assertRaises(TypeError): + DotMap(_default_factory='') + + def test_copy_preserves_default(self): + m = DotMap({'city': 'abc'}, _default_factory=str) + c = m.copy() + + self.assertEqual(c.city, 'abc') + self.assertNotIn('zipCode', c) + self.assertEqual(c.zipCode, '') + + def test_deepcopy_preserves_default(self): + m = DotMap({'address': {'city': 'abc'}}, _default_factory=str) + c = copy.deepcopy(m) + + self.assertEqual(c.address.city, 'abc') + self.assertNotIn('zipCode', c) + self.assertEqual(c.zipCode, '') + self.assertEqual(c.address.zipCode, '') + + def test_default_list(self): + m = DotMap({}, _default_factory=list) + self.assertEqual(m.x, []) + self.assertEqual(m.y, []) + + m.x.append(42) + self.assertEqual(m.x, [42]) + + self.assertEqual(m.y, []) + + def test_default_factory_returning_cached_list_shares_it(self): + # if the factory itself returns the same cached object every call, + # keys share it -- freshness is the factory's responsibility + cached = [] + m = DotMap({}, _default_factory=lambda: cached) + + self.assertIs(m.x, cached) + m.x.append(42) + + self.assertIs(m.y, cached) + self.assertEqual(m.y, [42]) + + class TestRecursive(unittest.TestCase): def test(self): m = DotMap()