[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:
@@ -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'])
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 """
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user