[FIX] *: strict api for main orm methods

Enforce strict types for returned values for
* create
* write
* unlink
* default_get

to make those methods more consistent and reliable.
Also make sure they can be called with empty self/values,
i.e. that they follow the same behavior as the base methods
defined in the main orm Model.

closes odoo/odoo#116809

Related: odoo/enterprise#38880
Signed-off-by: Victor Feyens (vfe) <vfe@odoo.com>
This commit is contained in:
Victor Feyens
2023-04-25 15:20:43 +02:00
parent 03aada713a
commit f4ea6d3226
27 changed files with 141 additions and 53 deletions
+6 -6
View File
@@ -871,10 +871,10 @@ class AccountGroup(models.Model):
@api.model_create_multi
def create(self, vals_list):
res_ids = super(AccountGroup, self).create([self._sanitize_vals(vals) for vals in vals_list])
res_ids._adapt_accounts_for_account_groups()
res_ids._adapt_parent_account_group()
return res_ids
groups = super().create([self._sanitize_vals(vals) for vals in vals_list])
groups._adapt_accounts_for_account_groups()
groups._adapt_parent_account_group()
return groups
def write(self, vals):
res = super(AccountGroup, self).write(self._sanitize_vals(vals))
@@ -890,7 +890,7 @@ class AccountGroup(models.Model):
children_ids = self.env['account.group'].search([('parent_id', '=', record.id)])
children_ids.write({'parent_id': record.parent_id.id})
super(AccountGroup, self).unlink()
return super().unlink()
def _adapt_accounts_for_account_groups(self, account_ids=None):
"""Ensure consistency between accounts and account groups.
@@ -903,7 +903,7 @@ class AccountGroup(models.Model):
return
company_ids = account_ids.company_id.ids if account_ids else self.company_id.ids
account_ids = account_ids.ids if account_ids else []
if not company_ids and account_ids is None:
if not company_ids and not account_ids:
return
self.flush_model()
self.env['account.account'].flush_model(['code'])
+1 -1
View File
@@ -1484,7 +1484,7 @@ class AccountMoveLine(models.Model):
def unlink(self):
if not self:
return
return True
# Check the lines are not reconciled (partially or not).
self._check_reconciliation()
+10 -7
View File
@@ -209,18 +209,21 @@ class AccountTax(models.Model):
@api.model
def default_get(self, fields_list):
# company_id is added so that we are sure to fetch a default value from it to use in repartition lines, below
rslt = super(AccountTax, self).default_get(fields_list + ['company_id'])
if 'company_id' not in fields_list and not {
'refund_repartition_line_ids',
'invoice_repartition_line_ids',
}.isdisjoint(fields_list):
fields_list += ['company_id']
rslt = super().default_get(fields_list)
company_id = rslt.get('company_id')
repartition = rslt.setdefault('repartition_line_ids', [])
if 'repartition_line_ids' in fields_list and not repartition:
repartition.extend([
if 'repartition_line_ids' in fields_list and 'repartition_line_ids' not in rslt:
company_id = rslt.get('company_id')
rslt['repartition_line_ids'] = [
Command.create({'document_type': 'invoice', 'repartition_type': 'base', 'tag_ids': [], 'company_id': company_id}),
Command.create({'document_type': 'invoice', 'repartition_type': 'tax', 'tag_ids': [], 'company_id': company_id}),
Command.create({'document_type': 'refund', 'repartition_type': 'base', 'tag_ids': [], 'company_id': company_id}),
Command.create({'document_type': 'refund', 'repartition_type': 'tax', 'tag_ids': [], 'company_id': company_id}),
])
]
return rslt
+2 -1
View File
@@ -61,10 +61,11 @@ class IrModule(models.Model):
def write(self, vals):
# Instanciate the first template of the module on the current company upon installing the module
was_installed = len(self) == 1 and self.state in ('installed', 'to upgrade', 'to remove')
super().write(vals)
res = super().write(vals)
is_installed = len(self) == 1 and self.state == 'installed'
if not was_installed and is_installed and not self.env.company.chart_template and self.account_templates:
self.env.registry._auto_install_template = next(iter(self.account_templates))
return res
def _load_module_terms(self, modules, langs, overwrite=False):
super()._load_module_terms(modules, langs, overwrite)
+1 -1
View File
@@ -28,4 +28,4 @@ class IrConfigParameter(models.Model):
if pls_emptied:
self.env.flush_all()
self.env.registry.setup_models(self.env.cr)
return pls_emptied
return result
+1 -1
View File
@@ -16,7 +16,7 @@ class Lead2OpportunityPartner(models.TransientModel):
to ease window action definitions, and be backward compatible. """
result = super(Lead2OpportunityPartner, self).default_get(fields)
if not result.get('lead_id') and self.env.context.get('active_id'):
if 'lead_id' in fields and not result.get('lead_id') and self.env.context.get('active_id'):
result['lead_id'] = self.env.context.get('active_id')
if result.get('lead_id'):
@@ -195,7 +195,7 @@ class DataRecycleModel(models.Model):
def write(self, vals):
if 'active' in vals and not vals['active']:
self.env['data_recycle.record'].search([('recycle_model_id', 'in', self.ids)]).unlink()
super().write(vals)
return super().write(vals)
def open_records(self):
self.ensure_one()
@@ -20,7 +20,8 @@ class DeliveryZipPrefix(models.Model):
return super().create(vals_list)
def write(self, vals):
vals['name'] = vals['name'].upper()
if 'name' in vals:
vals['name'] = vals['name'].upper()
return super().write(vals)
_sql_constraints = [
+2 -1
View File
@@ -306,8 +306,9 @@ class HrAttendance(models.Model):
def unlink(self):
attendances_dates = self._get_attendances_dates()
super(HrAttendance, self).unlink()
res = super().unlink()
self._update_overtime(attendances_dates)
return res
@api.returns('self', lambda value: value.id)
def copy(self, default=None):
+4 -2
View File
@@ -152,16 +152,18 @@ class LunchAlert(models.Model):
return alerts
def write(self, values):
super().write(values)
res = super().write(values)
if not CRON_DEPENDS.isdisjoint(values):
self._sync_cron()
return res
def unlink(self):
crons = self.cron_id.sudo()
server_actions = crons.ir_actions_server_id
super().unlink()
res = super().unlink()
crons.unlink()
server_actions.unlink()
return res
def _notify_chat(self):
# Called daily by cron
+4 -2
View File
@@ -199,16 +199,18 @@ class LunchSupplier(models.Model):
topping_values.update({'topping_category': 3})
if values.get('company_id'):
self.env['lunch.order'].search([('supplier_id', 'in', self.ids)]).write({'company_id': values['company_id']})
super().write(values)
res = super().write(values)
if not CRON_DEPENDS.isdisjoint(values):
self._sync_cron()
return res
def unlink(self):
crons = self.cron_id.sudo()
server_actions = crons.ir_actions_server_id
super().unlink()
res = super().unlink()
crons.unlink()
server_actions.unlink()
return res
def toggle_active(self):
""" Archiving related lunch product """
+3
View File
@@ -20,6 +20,9 @@ class IrModel(models.Model):
)
def unlink(self):
if not self:
return True
# Delete followers, messages and attachments for models that will be unlinked.
models = tuple(self.mapped('model'))
model_ids = tuple(self.ids)
+2 -2
View File
@@ -26,8 +26,8 @@ class MailActivity(models.Model):
@api.model
def default_get(self, fields):
res = super(MailActivity, self).default_get(fields)
if not fields or 'res_model_id' in fields and res.get('res_model'):
res = super().default_get(fields)
if 'res_model_id' in fields and res.get('res_model'):
res['res_model_id'] = self.env['ir.model']._get(res['res_model']).id
return res
+8 -5
View File
@@ -36,11 +36,14 @@ class MailBlackList(models.Model):
new_values.append(new_value)
""" To avoid crash during import due to unique email, return the existing records if any """
sql = '''SELECT email, id FROM mail_blacklist WHERE email = ANY(%s)'''
emails = [v['email'] for v in new_values]
self._cr.execute(sql, (emails,))
bl_entries = dict(self._cr.fetchall())
to_create = [v for v in new_values if v['email'] not in bl_entries]
to_create = []
bl_entries = {}
if new_values:
sql = '''SELECT email, id FROM mail_blacklist WHERE email = ANY(%s)'''
emails = [v['email'] for v in new_values]
self._cr.execute(sql, (emails,))
bl_entries = dict(self._cr.fetchall())
to_create = [v for v in new_values if v['email'] not in bl_entries]
# TODO DBE Fixme : reorder ids according to incoming ids.
results = super(MailBlackList, self).create(to_create)
+2 -2
View File
@@ -38,8 +38,8 @@ class MailGroup(models.Model):
@api.model
def default_get(self, fields):
res = super(MailGroup, self).default_get(fields)
if not res.get('alias_contact') and (not fields or 'alias_contact' in fields):
res = super().default_get(fields)
if 'alias_contact' in fields and not res.get('alias_contact'):
res['alias_contact'] = 'everyone' if res.get('access_mode') == 'public' else 'followers'
return res
@@ -41,11 +41,13 @@ class PhoneBlackList(models.Model):
to_create.append(dict(value, number=sanitized))
""" To avoid crash during import due to unique email, return the existing records if any """
sql = '''SELECT number, id FROM phone_blacklist WHERE number = ANY(%s)'''
numbers = [v['number'] for v in to_create]
self._cr.execute(sql, (numbers,))
bl_entries = dict(self._cr.fetchall())
to_create = [v for v in to_create if v['number'] not in bl_entries]
bl_entries = {}
if to_create:
sql = '''SELECT number, id FROM phone_blacklist WHERE number = ANY(%s)'''
numbers = [v['number'] for v in to_create]
self._cr.execute(sql, (numbers,))
bl_entries = dict(self._cr.fetchall())
to_create = [v for v in to_create if v['number'] not in bl_entries]
results = super(PhoneBlackList, self).create(to_create)
return self.env['phone.blacklist'].browse(bl_entries.values()) | results
+5 -4
View File
@@ -198,9 +198,10 @@ class PosSession(models.Model):
# installation we do the minimal configuration. Impossible to do in
# the .xml files as the CoA is not yet installed.
pos_config = self.env['pos.config'].browse(config_id)
ctx = dict(self.env.context, company_id=pos_config.company_id.id)
pos_name = self.env['ir.sequence'].with_context(ctx).next_by_code('pos.session')
pos_name = self.env['ir.sequence'].with_context(
company_id=pos_config.company_id.id
).next_by_code('pos.session')
if vals.get('name'):
pos_name += ' ' + vals['name']
@@ -213,9 +214,9 @@ class PosSession(models.Model):
})
if self.user_has_groups('point_of_sale.group_pos_user'):
sessions = super(PosSession, self.with_context(ctx).sudo()).create(vals_list)
sessions = super(PosSession, self.sudo()).create(vals_list)
else:
sessions = super(PosSession, self.with_context(ctx)).create(vals_list)
sessions = super().create(vals_list)
sessions.action_pos_session_open()
return sessions
+1 -1
View File
@@ -768,7 +768,7 @@ class Task(models.Model):
project = self.env['project.project'].browse(project_id)
if project.analytic_account_id:
vals['analytic_account_id'] = project.analytic_account_id.id
elif 'default_user_ids' not in self.env.context:
elif 'default_user_ids' not in self.env.context and 'user_ids' in default_fields:
user_ids = vals.get('user_ids', [])
user_ids.append(Command.link(self.env.user.id))
vals['user_ids'] = user_ids
+2 -2
View File
@@ -14,8 +14,8 @@ class SMSTemplate(models.Model):
@api.model
def default_get(self, fields):
res = super(SMSTemplate, self).default_get(fields)
if not fields or 'model_id' in fields and not res.get('model_id') and res.get('model'):
res = super().default_get(fields)
if 'model_id' in fields and not res.get('model_id') and res.get('model'):
res['model_id'] = self.env['ir.model']._get(res['model']).id
return res
+1 -1
View File
@@ -77,7 +77,7 @@ class UtmSourceMixin(models.AbstractModel):
if values.get('name'):
values['name'] = self.env['utm.mixin']._get_unique_names(self._name, [values['name']])[0]
super().write(values)
return super().write(values)
def copy(self, default=None):
"""Increment the counter when duplicating the source."""
+1 -1
View File
@@ -673,7 +673,7 @@ class Slide(models.Model):
def unlink(self):
for category in self.filtered(lambda slide: slide.is_category):
category.channel_id._move_category_slides(category, False)
super(Slide, self).unlink()
return super().unlink()
def toggle_active(self):
# archiving/unarchiving a channel does it on its slides, too
+2
View File
@@ -398,6 +398,8 @@ class ir_cron(models.Model):
the lock aquired by foreign keys when they
reference this row.
"""
if not self:
return
row_level_lock = "UPDATE" if lockfk else "NO KEY UPDATE"
try:
self._cr.execute(f"""
+6
View File
@@ -927,6 +927,9 @@ class IrModelFields(models.Model):
return res
def write(self, vals):
if not self:
return True
# if set, *one* column can be renamed here
column_rename = None
@@ -1404,6 +1407,9 @@ class IrModelSelection(models.Model):
return recs
def write(self, vals):
if not self:
return True
if (
not self.env.user._is_admin() and
any(record.field_id.state != 'manual' for record in self)
+4 -1
View File
@@ -21,6 +21,8 @@ def _create_sequence(cr, seq_name, number_increment, number_next):
def _drop_sequences(cr, seq_names):
""" Drop the PostreSQL sequences if they exist. """
if not seq_names:
return
names = sql.SQL(',').join(map(sql.Identifier, seq_names))
# RESTRICT is the default; it prevents dropping the sequence if an
# object depends on it.
@@ -335,7 +337,8 @@ class IrSequenceDateRange(models.Model):
@api.model
def default_get(self, fields):
result = super(IrSequenceDateRange, self).default_get(fields)
result['number_next_actual'] = 1
if 'number_next_actual' in fields:
result['number_next_actual'] = 1
return result
date_from = fields.Date(string='From', required=True)
+7 -5
View File
@@ -346,7 +346,7 @@ class ResConfigSettings(models.TransientModel, ResConfigModuleInstallationMixin)
The attribute 'group' may contain several xml ids, separated by commas.
* For a selection field like 'group_XXX' composed of 2 string values ('0' and '1'),
``execute`` adds/removes 'implied_group' to/from the implied groups of 'group',
``execute`` adds/removes 'implied_group' to/from the implied groups of 'group',
depending on the field's value.
By default 'group' is the group Employee. Groups are given by their xml id.
The attribute 'group' may contain several xml ids, separated by commas.
@@ -354,8 +354,8 @@ class ResConfigSettings(models.TransientModel, ResConfigModuleInstallationMixin)
* For a boolean field like 'module_XXX', ``execute`` triggers the immediate
installation of the module named 'XXX' if the field has value ``True``.
* For a selection field like 'module_XXX' composed of 2 string values ('0' and '1'),
``execute`` triggers the immediate installation of the module named 'XXX'
* For a selection field like 'module_XXX' composed of 2 string values ('0' and '1'),
``execute`` triggers the immediate installation of the module named 'XXX'
if the field has the value ``'1'``.
* For a field with no specific prefix BUT an attribute 'config_parameter',
@@ -466,12 +466,14 @@ class ResConfigSettings(models.TransientModel, ResConfigModuleInstallationMixin)
@api.model
def default_get(self, fields):
res = super().default_get(fields)
if not fields:
return res
IrDefault = self.env['ir.default']
IrConfigParameter = self.env['ir.config_parameter'].sudo()
classified = self._get_classified_fields(fields)
res = super(ResConfigSettings, self).default_get(fields)
# defaults: take the corresponding default value they set
for name, model, field in classified['default']:
value = IrDefault.get(model, field)
+1
View File
@@ -33,6 +33,7 @@ from . import test_module
from . import test_orm
from . import test_ormcache
from . import test_osv
from . import test_overrides
from . import test_qweb_field
from . import test_qweb
from . import test_res_config
+55
View File
@@ -0,0 +1,55 @@
# Part of Odoo. See LICENSE file for full copyright and licensing details.
from odoo.exceptions import UserError
from odoo.tests import TransactionCase, tagged
@tagged('-at_install', 'post_install')
class TestOverrides(TransactionCase):
# Ensure all main ORM methods behavior works fine even on empty recordset
# and that their returned value(s) follow the expected format.
def test_creates(self):
for model_env in self.env.values():
if model_env._abstract:
continue
# with self.assertQueryCount(0):
self.assertEqual(
model_env.create([]), model_env.browse(),
"Invalid create return value for model %s" % model_env._name)
def test_writes(self):
for model_env in self.env.values():
if model_env._abstract:
continue
try:
# with self.assertQueryCount(0):
self.assertEqual(
model_env.browse().write({}), True,
"Invalid write return value for model %s" % model_env._name)
except UserError:
# skip models that should never be modified
continue
def test_default_get(self):
for model_env in self.env.values():
if model_env._transient:
continue
try:
# with self.assertQueryCount(1): # allow one query for the call to get_model_defaults.
self.assertEqual(
model_env.browse().default_get([]), {},
"Invalid default_get return value for model %s" % model_env._name)
except UserError:
# skip "You must be logged in a Belgian company to use this feature" errors
continue
def test_unlink(self):
for model_env in self.env.values():
if model_env._abstract:
continue
# with self.assertQueryCount(0):
self.assertEqual(
model_env.browse().unlink(), True,
"Invalid unlink return value for model %s" % model_env._name)