[IMP] core: add named parameters to SQL wrapper

For very large bits of SQL code with potentially repeated terms, it is
useful to use named parameters instead of positional parameters:

    sql = SQL(
        "SELECT %(column)s FROM %(table)s WHERE %(column)s IS NOT NULL",
        table=SQL.identifier("foo"),
        column=SQL.identifier("foo", "bar"),
    )

Part-of: odoo/odoo#138019
This commit is contained in:
Raphael Collet
2023-10-20 09:35:17 +00:00
parent 0627e9940d
commit 31caff2834
4 changed files with 127 additions and 81 deletions
+22
View File
@@ -35,6 +35,12 @@ class TestSQL(BaseCase):
with self.assertRaises(TypeError):
SQL("SELECT id FROM table WHERE foo=%s AND bar=%s", 1, 2, 3)
with self.assertRaises(TypeError):
SQL("SELECT id FROM table WHERE foo=%s AND bar=%(two)s", 1, two=2)
with self.assertRaises(KeyError):
SQL("SELECT id FROM table WHERE foo=%(one)s AND bar=%(two)s", one=1, to=2)
def test_sql_equality(self):
sql1 = SQL("SELECT id FROM table WHERE foo=%s", 42)
sql2 = SQL("SELECT id FROM table WHERE foo=%s", 42)
@@ -115,6 +121,22 @@ class TestSQL(BaseCase):
self.assertEqual(sql.params, [1, 2])
self.assertEqual(sql, SQL("SELECT id FROM table WHERE foo=%s AND bar=%s", 1, 2))
def test_sql_with_named_parameters(self):
sql = SQL("SELECT id FROM table WHERE %(one)s AND bar=%(two)s", one=SQL("foo=%s", 1), two=2)
self.assertEqual(sql.code, "SELECT id FROM table WHERE foo=%s AND bar=%s")
self.assertEqual(sql.params, [1, 2])
self.assertEqual(sql, SQL("SELECT id FROM table WHERE foo=%s AND bar=%s", 1, 2))
# the parameters are bound locally
sql = SQL(
"%s AND %s",
SQL("foo=%(value)s", value=1),
SQL("bar=%(value)s", value=2),
)
self.assertEqual(sql.code, "foo=%s AND bar=%s")
self.assertEqual(sql.params, [1, 2])
self.assertEqual(sql, SQL("foo=%s AND bar=%s", 1, 2))
def test_complex_sql(self):
sql = SQL(
"SELECT %s FROM %s WHERE %s",
+61 -73
View File
@@ -2921,13 +2921,12 @@ class BaseModel(metaclass=MetaModel):
target_model = self.env[self._fields[field.definition_record].comodel_name]
self.env.cr.execute(SQL(
""" SELECT definition
FROM %s, jsonb_array_elements(%s) definition
WHERE %s IS NOT NULL AND definition->>'name' = %s
FROM %(table)s, jsonb_array_elements(%(field)s) definition
WHERE %(field)s IS NOT NULL AND definition->>'name' = %(name)s
LIMIT 1 """,
SQL.identifier(target_model._table),
SQL.identifier(field.definition_record_field),
SQL.identifier(field.definition_record_field),
property_name,
table=SQL.identifier(target_model._table),
field=SQL.identifier(field.definition_record_field),
name=property_name,
))
result = self.env.cr.dictfetchone()
return result["definition"] if result else {}
@@ -2955,21 +2954,18 @@ class BaseModel(metaclass=MetaModel):
query = SQL(
""" WITH RECURSIVE __parent_store_compute(id, parent_path) AS (
SELECT row.id, concat(row.id, '/')
FROM %s row
WHERE row.%s IS NULL
FROM %(table)s row
WHERE row.%(parent)s IS NULL
UNION
SELECT row.id, concat(comp.parent_path, row.id, '/')
FROM %s row, __parent_store_compute comp
WHERE row.%s = comp.id
FROM %(table)s row, __parent_store_compute comp
WHERE row.%(parent)s = comp.id
)
UPDATE %s row SET parent_path = comp.parent_path
UPDATE %(table)s row SET parent_path = comp.parent_path
FROM __parent_store_compute comp
WHERE row.id = comp.id """,
SQL.identifier(self._table),
SQL.identifier(self._parent_name),
SQL.identifier(self._table),
SQL.identifier(self._parent_name),
SQL.identifier(self._table),
table=SQL.identifier(self._table),
parent=SQL.identifier(self._parent_name),
)
self.env.cr.execute(query)
self.invalidate_model(['parent_path'])
@@ -3020,11 +3016,10 @@ class BaseModel(metaclass=MetaModel):
_logger.debug("Table '%s': setting default value of new column %s to %r",
self._table, column_name, value)
self._cr.execute(SQL(
"UPDATE %s SET %s = %s WHERE %s IS NULL",
SQL.identifier(self._table),
SQL.identifier(column_name),
value,
SQL.identifier(column_name),
"UPDATE %(table)s SET %(field)s = %(value)s WHERE %(field)s IS NULL",
table=SQL.identifier(self._table),
field=SQL.identifier(column_name),
value=value,
))
@ormcache()
@@ -3571,18 +3566,17 @@ class BaseModel(metaclass=MetaModel):
else next((v for v in translations.values() if v is not None), None)
self.invalidate_recordset([field_name])
self._cr.execute(SQL(
""" UPDATE %s
SET %s = NULLIF(
jsonb_strip_nulls(%s || COALESCE(%s, '{}'::jsonb) || %s),
""" UPDATE %(table)s
SET %(field)s = NULLIF(
jsonb_strip_nulls(%(fallback)s || COALESCE(%(field)s, '{}'::jsonb) || %(value)s),
'{}'::jsonb)
WHERE id = %s
WHERE id = %(id)s
""",
SQL.identifier(self._table),
SQL.identifier(field_name),
Json({'en_US': translation_fallback}),
SQL.identifier(field_name),
Json(translations),
self.id,
table=SQL.identifier(self._table),
field=SQL.identifier(field_name),
fallback=Json({'en_US': translation_fallback}),
value=Json(translations),
id=self.id,
))
self.modified([field_name])
else:
@@ -4433,11 +4427,10 @@ class BaseModel(metaclass=MetaModel):
# which fills the 'en_US' key of jsonb only when the old column value is NULL.
# The second param is for the real value {'fr_FR': 'French', 'nl_NL': 'Dutch'}
assignments.append(SQL(
"%s = %s || COALESCE(%s, '{}'::jsonb) || %s",
SQL.identifier(name),
Json({} if 'en_US' in val.adapted else {'en_US': next(iter(val.adapted.values()))}),
SQL.identifier(name),
val,
"%(field)s = %(fallback)s || COALESCE(%(field)s, '{}'::jsonb) || %(value)s",
field=SQL.identifier(name),
fallback=Json({} if 'en_US' in val.adapted else {'en_US': next(iter(val.adapted.values()))}),
value=val,
))
else:
assignments.append(SQL('%s = %s', SQL.identifier(name), val))
@@ -4830,18 +4823,17 @@ class BaseModel(metaclass=MetaModel):
return
self._cr.execute(SQL(
""" UPDATE %s node
""" UPDATE %(table)s node
SET parent_path=concat((
SELECT parent.parent_path
FROM %s parent
WHERE parent.id=node.%s
FROM %(table)s parent
WHERE parent.id=node.%(parent)s
), node.id, '/')
WHERE node.id IN %s
WHERE node.id IN %(ids)s
RETURNING node.id, node.parent_path """,
SQL.identifier(self._table),
SQL.identifier(self._table),
SQL.identifier(self._parent_name),
tuple(self.ids),
table=SQL.identifier(self._table),
parent=SQL.identifier(self._parent_name),
ids=tuple(self.ids),
))
# update the cache of updated nodes, and determine what to recompute
@@ -4860,15 +4852,14 @@ class BaseModel(metaclass=MetaModel):
parent_val = vals[self._parent_name]
if parent_val:
condition = SQL(
"(%s != %s OR %s IS NULL)",
SQL.identifier(self._parent_name),
parent_val,
SQL.identifier(self._parent_name),
"(%(parent)s != %(value)s OR %(parent)s IS NULL)",
parent=SQL.identifier(self._parent_name),
value=parent_val,
)
else:
condition = SQL(
"%s IS NOT NULL",
SQL.identifier(self._parent_name),
"%(parent)s IS NOT NULL",
parent=SQL.identifier(self._parent_name),
)
self._cr.execute(SQL(
"SELECT id FROM %s WHERE id IN %s AND %s",
@@ -4885,12 +4876,11 @@ class BaseModel(metaclass=MetaModel):
# determine new prefix of parent_path
cr.execute(SQL(
""" SELECT parent.parent_path
FROM %s node, %s parent
WHERE node.id = %s AND parent.id = node.%s """,
SQL.identifier(self._table),
SQL.identifier(self._table),
self.ids[0],
SQL.identifier(self._parent_name),
FROM %(table)s node, %(table)s parent
WHERE node.id = %(id)s AND parent.id = node.%(parent)s """,
table=SQL.identifier(self._table),
parent=SQL.identifier(self._parent_name),
id=self.ids[0],
))
prefix = cr.fetchone()[0] if cr.rowcount else ''
@@ -4902,18 +4892,17 @@ class BaseModel(metaclass=MetaModel):
# update parent_path of all records and their descendants
cr.execute(SQL(
""" UPDATE %s child
SET parent_path = concat(%s, substr(child.parent_path,
""" UPDATE %(table)s child
SET parent_path = concat(%(prefix)s, substr(child.parent_path,
length(node.parent_path) - length(node.id || '/') + 1))
FROM %s node
WHERE node.id IN %s
AND child.parent_path LIKE concat(node.parent_path, %s)
FROM %(table)s node
WHERE node.id IN %(ids)s
AND child.parent_path LIKE concat(node.parent_path, %(wildcard)s)
RETURNING child.id, child.parent_path """,
SQL.identifier(self._table),
prefix,
SQL.identifier(self._table),
tuple(self.ids),
'%',
table=SQL.identifier(self._table),
prefix=prefix,
ids=tuple(self.ids),
wildcard='%',
))
# update the cache of updated nodes, and determine what to recompute
@@ -5560,13 +5549,12 @@ class BaseModel(metaclass=MetaModel):
while todo:
# retrieve the respective successors of the nodes in 'todo'
cr.execute(SQL(
"SELECT %s, %s FROM %s WHERE %s IN %s AND %s IS NOT NULL",
SQL.identifier(field.column1),
SQL.identifier(field.column2),
SQL.identifier(field.relation),
SQL.identifier(field.column1),
tuple(todo),
SQL.identifier(field.column2),
""" SELECT %(col1)s, %(col2)s FROM %(rel)s
WHERE %(col1)s IN %(ids)s AND %(col2)s IS NOT NULL """,
rel=SQL.identifier(field.relation),
col1=SQL.identifier(field.column1),
col2=SQL.identifier(field.column2),
ids=tuple(todo),
))
done.update(todo)
todo.clear()
+24
View File
@@ -1824,3 +1824,27 @@ def get_flag(country_code: str) -> str:
def format_frame(frame):
code = frame.f_code
return f'{code.co_name} {code.co_filename}:{frame.f_lineno}'
def named_to_positional_printf(string: str, args: Mapping) -> tuple[str, tuple]:
""" Convert a named printf-style format string with its arguments to an
equivalent positional format string with its arguments. This implementation
does not support escaped ``%`` characters (``"%%"``).
"""
if '%%' in string:
raise ValueError(f"Unsupported escaped '%' in format string {string!r}")
args = _PrintfArgs(args)
return string % args, tuple(args.values)
class _PrintfArgs:
""" Helper object to turn a named printf-style format string into a positional one. """
__slots__ = ('mapping', 'values')
def __init__(self, mapping):
self.mapping = mapping
self.values = []
def __getitem__(self, key):
self.values.append(self.mapping[key])
return "%s"
+20 -8
View File
@@ -13,6 +13,8 @@ from typing import Iterable, Union
import psycopg2
from .misc import named_to_positional_printf
_schema = logging.getLogger('odoo.schema')
IDENT_RE = re.compile(r'^[a-z0-9_][a-z0-9_$\-]*$', re.I)
@@ -27,15 +29,18 @@ _CONFDELTYPES = {
class SQL:
""" An object that wraps SQL code with its positional parameters, like::
""" An object that wraps SQL code with its parameters, like::
sql = SQL("UPDATE TABLE foo SET a = %s, b = %s", 'hello', 42)
cr.execute(sql)
The code is given as a format string, and the positional arguments are meant
to be merged into it using the string formatting operator. The wrapper is
designed to be composable: the positional parameters can be either actual
parameters or SQL objects themselves::
The code is given as a ``%``-format string, and supports either positional
arguments (with `%s`) or named arguments (with `%(name)s`). Escaped
characters (like ``"%%"``) are not supported, though. The arguments are
meant to be merged into the code using the `%` formatting operator.
The SQL wrapper is designed to be composable: the arguments can be either
actual parameters, or SQL objects themselves::
sql = SQL(
"UPDATE TABLE %s SET %s",
@@ -57,11 +62,18 @@ class SQL:
__slots__ = ('__code', '__args')
# pylint: disable=keyword-arg-before-vararg
def __new__(cls, code: (str | SQL) = "", /, *args):
def __new__(cls, code: (str | SQL) = "", /, *args, **kwargs):
if isinstance(code, SQL):
return code
# validate the format of code
code % tuple("" for arg in args)
# validate the format of code and parameters
if args and kwargs:
raise TypeError("SQL() takes either positional arguments, or named arguments")
if args:
code % tuple("" for arg in args)
elif kwargs:
code, args = named_to_positional_printf(code, kwargs)
self = object.__new__(cls)
self.__code = code
self.__args = args