Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
50 changes: 41 additions & 9 deletions dotmap/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -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:
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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 = '['
Expand Down
65 changes: 65 additions & 0 deletions dotmap/test.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import unittest
import copy
from dotmap import DotMap


Expand Down Expand Up @@ -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}
Comment thread
sachinsachdeva marked this conversation as resolved.
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()
Expand Down