From 32a58c0db3c0d666e704a4c8af33b0be0a9608e4 Mon Sep 17 00:00:00 2001 From: Raphael Collet Date: Mon, 28 Aug 2017 14:30:38 +0200 Subject: [PATCH] [REF] api: wrap the record cache implementation into a class This makes the rest of the ORM independent from the cache's implementation, and rely on a well-defined API instead. --- odoo/addons/base/tests/test_api.py | 37 ++-- .../test_new_api/tests/test_onchange.py | 30 +-- odoo/api.py | 171 ++++++++++++------ odoo/fields.py | 62 ++++--- odoo/models.py | 62 +++---- 5 files changed, 205 insertions(+), 157 deletions(-) diff --git a/odoo/addons/base/tests/test_api.py b/odoo/addons/base/tests/test_api.py index 96fba1179b7..90a2987b8da 100644 --- a/odoo/addons/base/tests/test_api.py +++ b/odoo/addons/base/tests/test_api.py @@ -250,11 +250,11 @@ class TestAPI(common.TransactionCase): # fetch data in the cache for p in partners: p.name, p.company_id.name, p.user_id.name, p.contact_address - self.env.check_cache() + self.env.cache.check(self.env) # change its parent child.write({'parent_id': partner2.id}) - self.env.check_cache() + self.env.cache.check(self.env) # check recordsets self.assertEqual(child.parent_id, partner2) @@ -262,16 +262,16 @@ class TestAPI(common.TransactionCase): self.assertIn(child, partner2.child_ids) self.assertEqual(set(partner1.child_ids + child), set(children1)) self.assertEqual(set(partner2.child_ids), set(children2 + child)) - self.env.check_cache() + self.env.cache.check(self.env) # delete it child.unlink() - self.env.check_cache() + self.env.cache.check(self.env) # check recordsets self.assertEqual(set(partner1.child_ids), set(children1) - set([child])) self.assertEqual(set(partner2.child_ids), set(children2)) - self.env.check_cache() + self.env.cache.check(self.env) # convert from the cache format to the write format partner = partner1 @@ -290,21 +290,30 @@ class TestAPI(common.TransactionCase): self.assertItemsEqual(partners.ids, partners._prefetch['res.partner']) # reading ONE partner should fetch them ALL - partner = next(p for p in partners) - partner.country_id - country_id_cache = self.env.cache[type(partners).country_id] - self.assertItemsEqual(partners.ids, country_id_cache) + for partner in partners: + partner.country_id + break + partner_ids_with_field = [partner.id + for partner in partners + if 'country_id' in partner._cache] + self.assertItemsEqual(partner_ids_with_field, partners.ids) # partners' countries are ready for prefetching - country_ids = set(cid for cids in country_id_cache.values() for cid in cids) + country_ids = {cid + for partner in partners + for cid in partner._cache['country_id']} self.assertTrue(len(country_ids) > 1) self.assertItemsEqual(country_ids, partners._prefetch['res.country']) # reading ONE partner country should fetch ALL partners' countries - country = next(p.country_id for p in partners if p.country_id) - country.name - name_cache = self.env.cache[type(country).name] - self.assertItemsEqual(country_ids, name_cache) + for partner in partners: + if partner.country_id: + partner.country_id.name + break + country_ids_with_field = [country.id + for country in partners.mapped('country_id') + if 'name' in country._cache] + self.assertItemsEqual(country_ids_with_field, country_ids) @mute_logger('odoo.models') def test_60_prefetch_object(self): diff --git a/odoo/addons/test_new_api/tests/test_onchange.py b/odoo/addons/test_new_api/tests/test_onchange.py index f35763bafd8..e70b1736b30 100644 --- a/odoo/addons/test_new_api/tests/test_onchange.py +++ b/odoo/addons/test_new_api/tests/test_onchange.py @@ -44,7 +44,7 @@ class TestOnChange(common.TransactionCase): 'author': USER.id, 'size': 0, } - self.env.invalidate_all() + self.env.cache.invalidate() result = self.Message.onchange(values, 'discussion', field_onchange) self.assertIn('name', result['value']) self.assertEqual(result['value']['name'], "[%s] %s" % (discussion.name, USER.name)) @@ -57,7 +57,7 @@ class TestOnChange(common.TransactionCase): 'author': USER.id, 'size': 0, } - self.env.invalidate_all() + self.env.cache.invalidate() result = self.Message.onchange(values, 'body', field_onchange) self.assertIn('size', result['value']) self.assertEqual(result['value']['size'], len(BODY)) @@ -71,7 +71,7 @@ class TestOnChange(common.TransactionCase): 'author': USER.id, 'size': 0, } - self.env.invalidate_all() + self.env.cache.invalidate() result = self.Message.onchange(values, 'body', field_onchange) self.assertNotIn('name', result['value']) @@ -89,7 +89,7 @@ class TestOnChange(common.TransactionCase): 'root_categ': False, } - self.env.invalidate_all() + self.env.cache.invalidate() result = Category.onchange(values, 'parent', field_onchange).get('value', {}) self.assertIn('root_categ', result) self.assertEqual(result['root_categ'], root.name_get()[0]) @@ -97,7 +97,7 @@ class TestOnChange(common.TransactionCase): values.update(result) values['parent'] = False - self.env.invalidate_all() + self.env.cache.invalidate() result = Category.onchange(values, 'parent', field_onchange).get('value', {}) self.assertIn('root_categ', result) self.assertIs(result['root_categ'], False) @@ -136,7 +136,7 @@ class TestOnChange(common.TransactionCase): }), ], } - self.env.invalidate_all() + self.env.cache.invalidate() result = self.Discussion.onchange(values, 'name', field_onchange) self.assertIn('messages', result['value']) self.assertItemsEqual(result['value']['messages'], [ @@ -188,7 +188,7 @@ class TestOnChange(common.TransactionCase): }), ], } - self.env.invalidate_all() + self.env.cache.invalidate() result = self.Discussion.onchange(values, 'name', field_onchange) self.assertIn('messages', result['value']) self.assertItemsEqual(result['value']['messages'], [ @@ -231,7 +231,7 @@ class TestOnChange(common.TransactionCase): partner = self.env.ref('base.res_partner_2') values['partner'] = partner.id values['lines'].append((0, 0, {'name': False, 'partner': False})) - self.env.invalidate_all() + self.env.cache.invalidate() result = multi.onchange(values, 'partner', field_onchange) self.assertEqual(result['value'], { 'name': partner.name, @@ -266,7 +266,7 @@ class TestOnChange(common.TransactionCase): 'messages': [(4, msg.id) for msg in discussion.messages], 'participants': [(4, usr.id) for usr in discussion.participants], } - self.env.invalidate_all() + self.env.cache.invalidate() result = discussion.onchange(values, 'moderator', field_onchange) self.assertIn('participants', result['value']) @@ -287,13 +287,13 @@ class TestOnChange(common.TransactionCase): self.env['ir.default'].set('test_new_api.foo', 'value2', 666, condition='value1=42') # setting 'value1' to 42 should trigger the change of 'value2' - self.env.invalidate_all() + self.env.cache.invalidate() values = {'name': 'X', 'value1': 42, 'value2': False} result = Foo.onchange(values, 'value1', field_onchange) self.assertEqual(result['value'], {'value2': 666}) # setting 'value1' to 24 should not trigger the change of 'value2' - self.env.invalidate_all() + self.env.cache.invalidate() values = {'name': 'X', 'value1': 24, 'value2': False} result = Foo.onchange(values, 'value1', field_onchange) self.assertEqual(result['value'], {}) @@ -349,16 +349,16 @@ class TestOnChange(common.TransactionCase): }) # check if server-side cache is working correctly - self.env.invalidate_all() + self.env.cache.invalidate() self.assertIn(email, discussion.emails) self.assertNotIn(email, discussion.important_emails) email.important = True self.assertIn(email, discussion.important_emails) # check that when trigger an onchange, we don't reset important emails - # (force `invalidate_all` as but appear in onchange only when we get a - # cache miss) - self.env.invalidate_all() + # (force `invalidate` as but appear in onchange only when we get a cache + # miss) + self.env.cache.invalidate() self.assertEqual(len(discussion.messages), 4) values = { 'name': "Foo Bar", diff --git a/odoo/api.py b/odoo/api.py index dbac3921988..89f58b613b7 100644 --- a/odoo/api.py +++ b/odoo/api.py @@ -740,7 +740,7 @@ class Environment(Mapping): self = object.__new__(cls) self.cr, self.uid, self.context = self.args = (cr, uid, frozendict(context)) self.registry = Registry(cr.dbname) - self.cache = defaultdict(dict) # {field: {id: value, ...}, ...} + self.cache = Cache() self._protected = defaultdict(frozenset) # {field: ids, ...} self.dirty = defaultdict(set) # {record: set(field_name), ...} self.all = envs @@ -836,39 +836,11 @@ class Environment(Mapping): """ Return whether we are in 'onchange' draft mode. """ return self.all.mode == 'onchange' - def invalidate(self, spec): - """ Invalidate some fields for some records in the cache of all - environments. - - :param spec: what to invalidate, a list of `(field, ids)` pair, - where ``field`` is a field object, and ``ids`` is a list of record - ids or ``None`` (to invalidate all records). - """ - if not spec: - return - for env in list(self.all): - c = env.cache - for field, ids in spec: - if ids is None: - if field in c: - del c[field] - else: - field_cache = c[field] - for id in ids: - field_cache.pop(id, None) - - def invalidate_all(self): - """ Clear the cache of all environments. """ - for env in list(self.all): - env.cache.clear() - env._protected.clear() - env.dirty.clear() - def clear(self): """ Clear all record caches, and discard all fields to recompute. This may be useful when recovering from a failed ORM operation. """ - self.invalidate_all() + self.cache.invalidate() self.all.todo.clear() @contextmanager @@ -939,36 +911,6 @@ class Environment(Mapping): field = min(self.all.todo, key=self.registry.field_sequence) return field, self.all.todo[field][0] - def check_cache(self): - """ Check the cache consistency. """ - from odoo.models import SpecialValue - - # make a full copy of the cache, and invalidate it - cache_dump = dict( - (field, dict(field_cache)) - for field, field_cache in self.cache.items() - ) - self.invalidate_all() - - # re-fetch the records, and compare with their former cache - invalids = [] - for field, field_dump in cache_dump.items(): - records = self[field.model_name].browse(f for f in field_dump if f) - for record in records: - try: - cached = field_dump[record.id] - cached = cached.get() if isinstance(cached, SpecialValue) else cached - value = field.convert_to_record(cached, record) - fetched = record[field.name] - if fetched != value: - info = {'cached': value, 'fetched': fetched} - invalids.append((field, record, info)) - except (AccessError, MissingError): - pass - - if invalids: - raise UserError('Invalid cache for fields\n' + pformat(invalids)) - @property def recompute(self): return self.all.recompute @@ -1000,6 +942,115 @@ class Environments(object): return iter(self.envs) +class Cache(object): + """ Implementation of the cache of records. """ + def __init__(self): + self._data = defaultdict(dict) # {field: {id: value}} + + def contains(self, record, field): + """ Return whether ``record`` has a value for ``field``. """ + return record.id in self._data[field] + + def get(self, record, field): + """ Return the value of ``field`` for ``record``. """ + value = self._data[field][record.id] + return value.get() if isinstance(value, SpecialValue) else value + + def set(self, record, field, value): + """ Set the value of ``field`` for ``record``. """ + self._data[field][record.id] = value + + def remove(self, record, field): + """ Remove the value of ``field`` for ``record``. """ + del self._data[field][record.id] + + def contains_value(self, record, field): + """ Return whether ``record`` has a regular value for ``field``. """ + value = self._data[field].get(record.id, SpecialValue(None)) + return not isinstance(value, SpecialValue) + + def get_value(self, record, field, default=None): + """ Return the regular value of ``field`` for ``record``. """ + value = self._data[field].get(record.id, SpecialValue(None)) + return default if isinstance(value, SpecialValue) else value + + def set_special(self, record, field, getter): + """ Set the value of ``field`` for ``record`` to return ``getter()``. """ + self._data[field][record.id] = SpecialValue(getter) + + def set_failed(self, records, fields, exception): + """ Mark ``fields`` on ``records`` with the given exception. """ + def getter(): + raise exception + for field in fields: + for record in records: + self.set_special(record, field, getter) + + def get_fields(self, record): + """ Return the fields with a value for ``record``. """ + for name, field in record._fields.items(): + if name != 'id' and record.id in self._data[field]: + yield field + + def get_records(self, model, field): + """ Return the records of ``model`` that have a value for ``field``. """ + return model.browse(self._data[field]) + + def invalidate(self, spec=None): + """ Invalidate the cache, partially or totally depending on ``spec``. """ + if spec is None: + for env in list(Environment.envs): + env.cache._data.clear() + elif spec: + for env in list(Environment.envs): + data = env.cache._data + for field, ids in spec: + if ids is None: + if field in data: + del data[field] + else: + field_cache = data[field] + for id in ids: + field_cache.pop(id, None) + + def check(self, env): + """ Check the consistency of the cache for the given environment. """ + # make a full copy of the cache, and invalidate it + dump = defaultdict(dict) + for field, field_cache in env.cache._data.items(): + for record_id, value in field_cache.items(): + if record_id: + dump[field][record_id] = value + self.invalidate() + + # re-fetch the records, and compare with their former cache + invalids = [] + for field, field_dump in dump.items(): + records = env[field.model_name].browse(field_dump) + for record in records: + try: + cached = field_dump[record.id] + cached = cached.get() if isinstance(cached, SpecialValue) else cached + value = field.convert_to_record(cached, record) + fetched = record[field.name] + if fetched != value: + info = {'cached': value, 'fetched': fetched} + invalids.append((record, field, info)) + except (AccessError, MissingError): + pass + + if invalids: + raise UserError('Invalid cache for fields\n' + pformat(invalids)) + + +class SpecialValue(object): + """ Wrapper for a function to get the cached value of a field. """ + __slots__ = ['get'] + + def __init__(self, getter): + self.get = getter + + # keep those imports here in order to handle cyclic dependencies correctly from odoo import SUPERUSER_ID from odoo.exceptions import UserError, AccessError, MissingError diff --git a/odoo/fields.py b/odoo/fields.py index 9a2d737b7f2..f8b3efb486c 100644 --- a/odoo/fields.py +++ b/odoo/fields.py @@ -39,18 +39,18 @@ Default = object() # default value for __init__() methods def copy_cache(records, env): """ Recursively copy the cache of ``records`` to the environment ``env``. """ + src, dst = records.env.cache, env.cache todo, done = set(records), set() while todo: record = todo.pop() if record not in done: done.add(record) target = record.with_env(env) - for name in record._cache: - field = record._fields[name] - value = record[name] - if isinstance(value, BaseModel): - todo.update(value) - target._cache[name] = field.convert_to_cache(value, target, validate=False) + for field in src.get_fields(record): + value = src.get(record, field) + dst.set(target, field, value) + if value and field.type in ('many2one', 'one2many', 'many2many', 'reference'): + todo.update(field.convert_to_record(value, record)) def resolve_mro(model, name, predicate): @@ -893,14 +893,14 @@ class Field(MetaField('DummyField', (object,), {})): # only a single record may be accessed record.ensure_one() try: - value = record._cache[self.name] + value = record.env.cache.get(record, self) except KeyError: # cache miss, determine value and retrieve it if record.id: self.determine_value(record) else: self.determine_draft_value(record) - value = record._cache[self.name] + value = record.env.cache.get(record, self) else: # null record -> return the null value for this field value = self.convert_to_cache(False, record, validate=False) @@ -922,7 +922,7 @@ class Field(MetaField('DummyField', (object,), {})): spec = self.modified_draft(record) # set value in cache, inverse field, and mark record as dirty - record._cache[self.name] = value + record.env.cache.set(record, self, value) if env.in_onchange: for invf in record._field_inverses[self]: invf._update(record[self.name], record) @@ -931,7 +931,7 @@ class Field(MetaField('DummyField', (object,), {})): # determine more dependent fields, and invalidate them if self.relational: spec += self.modified_draft(record) - env.invalidate(spec) + env.cache.invalidate(spec) else: # Write to database @@ -939,7 +939,7 @@ class Field(MetaField('DummyField', (object,), {})): record.write({self.name: write_value}) # Update the cache unless value contains a new record if not (self.relational and not all(value)): - record._cache[self.name] = value + record.env.cache.set(record, self, value) ############################################################################ # @@ -950,9 +950,10 @@ class Field(MetaField('DummyField', (object,), {})): """ Invoke the compute method on ``records``. """ # initialize the fields to their corresponding null value in cache fields = records._field_computed[self] + cache = records.env.cache for field in fields: for record in records: - record._cache[field.name] = field.convert_to_cache(False, record, validate=False) + cache.set(record, field, field.convert_to_cache(False, record, validate=False)) if isinstance(self.compute, pycompat.string_types): getattr(records, self.compute)() else: @@ -970,7 +971,7 @@ class Field(MetaField('DummyField', (object,), {})): try: self._compute_value(record) except Exception as exc: - record._cache.set_failed([self.name], exc) + record.env.cache.set_failed(record, [self], exc) def determine_value(self, record): """ Determine the value of ``self`` for ``record``. """ @@ -1011,7 +1012,7 @@ class Field(MetaField('DummyField', (object,), {})): else: # this is a non-stored non-computed field - record._cache[self.name] = self.convert_to_cache(False, record, validate=False) + record.env.cache.set(record, self, self.convert_to_cache(False, record, validate=False)) def determine_draft_value(self, record): """ Determine the value of ``self`` for the given draft ``record``. """ @@ -1021,7 +1022,7 @@ class Field(MetaField('DummyField', (object,), {})): self._compute_value(record) else: null = self.convert_to_cache(False, record, validate=False) - record._cache.set_special(self.name, lambda: null) + record.env.cache.set_special(record, self, lambda: null) def determine_inverse(self, records): """ Given the value of ``self`` on ``records``, inverse the computation. """ @@ -1062,11 +1063,11 @@ class Field(MetaField('DummyField', (object,), {})): if path == 'id' and field.model_name == records._name: target = records - protected elif path and env.in_onchange: - target = (target.browse(env.cache[field]) - protected).filtered( + target = (env.cache.get_records(target, field) - protected).filtered( lambda rec: rec if path == 'id' else rec._mapped_cache(path) & records ) else: - target = target.browse(env.cache[field]) - protected + target = env.cache.get_records(target, field) - protected if target: spec.append((field, target._ids)) @@ -1117,8 +1118,9 @@ class Integer(Field): def _update(self, records, value): # special case, when an integer field is used as inverse for a one2many + cache = records.env.cache for record in records: - record._cache[self.name] = value.id or 0 + cache.set(record, self, value.id or 0) def convert_to_export(self, value, record): if value or value == 0: @@ -1632,8 +1634,9 @@ class Binary(Field): # Note: the 'bin_size' flag is handled by the field 'datas' itself data = {att.res_id: att.datas for att in records.env['ir.attachment'].sudo().search(domain)} + cache = records.env.cache for record in records: - record._cache[self.name] = data.get(record.id, False) + cache.set(record, self, data.get(record.id, False)) def write(self, records, value): # retrieve the attachments that stores the value, and adapt them @@ -1905,8 +1908,9 @@ class Many2one(_Relational): def _update(self, records, value): """ Update the cached value of ``self`` for ``records`` with ``value``. """ + cache = records.env.cache for record in records: - record._cache[self.name] = self.convert_to_cache(value, record, validate=False) + cache.set(record, self, self.convert_to_cache(value, record, validate=False)) def convert_to_column(self, value, record, values=None): return value or None @@ -1967,19 +1971,21 @@ class _RelationalMulti(_Relational): def _update(self, records, value): """ Update the cached value of ``self`` for ``records`` with ``value``. """ + cache = records.env.cache for record in records: - if self.name in record._cache: + if cache.contains(record, self): val = self.convert_to_cache(record[self.name] | value, record, validate=False) - record._cache[self.name] = val + cache.set(record, self, val) else: - record._cache.set_special(self.name, self._update_getter(record, value)) + cache.set_special(record, self, self._update_getter(record, value)) def _update_getter(self, record, value): def getter(): # determine the current field's value, and update it in cache only - del record._cache[self.name] + cache = record.env.cache + cache.remove(record, self) val = self.convert_to_cache(record[self.name] | value, record, validate=False) - record._cache[self.name] = val + cache.set(record, self, val) return val return getter @@ -2173,8 +2179,9 @@ class One2many(_RelationalMulti): group[int(line[inverse])].append(line.id) # store result in cache + cache = records.env.cache for record in records: - record._cache[self.name] = tuple(group[record.id]) + cache.set(record, self, tuple(group[record.id])) def write(self, records, value): comodel = records.env[self.comodel_name].with_context(**self.context) @@ -2370,8 +2377,9 @@ class Many2many(_RelationalMulti): group[row[0]].append(row[1]) # store result in cache + cache = records.env.cache for record in records: - record._cache[self.name] = tuple(group[record.id]) + cache.set(record, self, tuple(group[record.id])) def write(self, records, value): cr = records._cr diff --git a/odoo/models.py b/odoo/models.py index 31a096673f1..9c60a0da2bc 100644 --- a/odoo/models.py +++ b/odoo/models.py @@ -2555,7 +2555,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): # in onchange mode, discard computed fields and fields in cache if self.env.in_onchange: for f in list(fs): - if f.compute or (f.name in self._cache): + if f.compute or self.env.cache.contains(self, f): fs.discard(f) else: records &= self._in_cache_without(f) @@ -2571,13 +2571,13 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): result = self.read([f.name for f in fs], load='_classic_write') # check the cache, and update it if necessary - if not self._cache.has_value(field.name): + if not self.env.cache.contains_value(self, field): for values in result: record = self.browse(values.pop('id'), self._prefetch) record._cache.update(record._convert_to_cache(values, validate=False)) - if field.name not in self._cache: + if not self.env.cache.contains(self, field): exc = AccessError("No value found for %s.%s" % (self, field.name)) - self._cache.set_failed([field.name], exc) + self.env.cache.set_failed(self, field, exc) @api.multi def _read_from_database(self, field_names, inherited_field_names=[]): @@ -2682,8 +2682,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): _('The requested operation cannot be completed due to security restrictions. Please contact your system administrator.\n\n(Document type: %s, Operation: %s)') % \ (self._name, 'read') ) - for record in forbidden: - record._cache.set_failed(self._fields, exc) + self.env.cache.set_failed(forbidden, self._fields.values(), exc) @api.multi def get_metadata(self): @@ -3870,8 +3869,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): if len(existing) < len(self): # mark missing records in cache with a failed value exc = MissingError(_("Record does not exist or has been deleted.")) - for record in (self - existing): - record._cache.set_failed(self._fields, exc) + self.env.cache.set_failed(self - existing, self._fields.values(), exc) return existing @api.multi @@ -4680,8 +4678,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): (:class:`Field` instance), including ``self``. Return at most ``limit`` records. """ - ids = [it for it in self._prefetch[self._name] - set(self.env.cache[field]) if it] - recs = self.browse(ids) + recs = self.browse(self._prefetch[self._name]) - self.env.cache.get_records(self, field) if limit and len(recs) > limit: recs = self + (recs - self)[:(limit - len(self))] return recs @@ -4705,7 +4702,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): """ if fnames is None: if ids is None: - return self.env.invalidate_all() + return self.env.cache.invalidate() fields = list(self._fields.values()) else: fields = [self._fields[n] for n in fnames] @@ -4713,7 +4710,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): # invalidate fields and inverse fields, too spec = [(f, ids) for f in fields] + \ [(invf, None) for f in fields for invf in self._field_inverses[f]] - self.env.invalidate(spec) + self.env.cache.invalidate(spec) @api.multi def modified(self, fnames): @@ -4765,7 +4762,7 @@ class BaseModel(MetaModel('DummyModel', (object,), {'_register': False})): for field in (fields - stored): invalids.append((field, None)) - self.env.invalidate(invalids) + self.env.cache.invalidate(invalids) def _recompute_check(self, field): """ If ``field`` must be recomputed on some record in ``self``, return the @@ -5024,30 +5021,27 @@ class RecordCache(MutableMapping): def __contains__(self, name): """ Return whether `record` has a cached value for field ``name``. """ field = self._record._fields[name] - return self._record.id in self._record.env.cache[field] + return self._record.env.cache.contains(self._record, field) def __getitem__(self, name): """ Return the cached value of field ``name`` for `record`. """ field = self._record._fields[name] - value = self._record.env.cache[field][self._record.id] - return value.get() if isinstance(value, SpecialValue) else value + return self._record.env.cache.get(self._record, field) def __setitem__(self, name, value): """ Assign the cached value of field ``name`` for ``record``. """ field = self._record._fields[name] - self._record.env.cache[field][self._record.id] = value + self._record.env.cache.set(self._record, field, value) def __delitem__(self, name): """ Remove the cached value of field ``name`` for ``record``. """ field = self._record._fields[name] - del self._record.env.cache[field][self._record.id] + self._record.env.cache.remove(self._record, field) def __iter__(self): """ Iterate over the field names with a cached value. """ - cache, record_id = self._record.env.cache, self._record.id - for name, field in self._record._fields.items(): - if name != 'id' and record_id in cache[field]: - yield name + for field in self._record.env.cache.get_fields(self._record): + yield field.name def __len__(self): """ Return the number of fields with a cached value. """ @@ -5056,36 +5050,22 @@ class RecordCache(MutableMapping): def has_value(self, name): """ Return whether `record` has a cached, regular value for field ``name``. """ field = self._record._fields[name] - dummy = SpecialValue(None) - value = self._record.env.cache[field].get(self._record.id, dummy) - return not isinstance(value, SpecialValue) + return self._record.env.cache.contains_value(self._record, field) def get_value(self, name, default=None): """ Return the cached, regular value of field ``name`` for `record`, or ``default``. """ field = self._record._fields[name] - dummy = SpecialValue(None) - value = self._record.env.cache[field].get(self._record.id, dummy) - return default if isinstance(value, SpecialValue) else value + return self._record.env.cache.get_value(self._record, field, default) def set_special(self, name, getter): """ Use the given getter to get the cached value of field ``name``. """ field = self._record._fields[name] - self._record.env.cache[field][self._record.id] = SpecialValue(getter) + self._record.env.cache.set_special(self._record, field, getter) def set_failed(self, names, exception): """ Mark the given fields with the given exception. """ - def getter(): - raise exception - for name in names: - self.set_special(name, getter) - - -class SpecialValue(object): - """ Wrapper for a function to get the cached value of a field. """ - __slots__ = ['get'] - - def __init__(self, getter): - self.get = getter + fields = [self._record._fields[name] for name in names] + self._record.env.cache.set_failed(self._record, fields, exception) AbstractModel = BaseModel