[IMP] core: type annotations in SQL wrapper and Query

Part-of: odoo/odoo#138019
This commit is contained in:
Raphael Collet
2023-10-20 09:35:17 +00:00
parent d68696e11e
commit 0627e9940d
2 changed files with 14 additions and 15 deletions
+9 -11
View File
@@ -1,26 +1,24 @@
# -*- coding: utf-8 -*-
# Part of Odoo. See LICENSE file for full copyright and licensing details.
from typing import Union
from odoo.tools.sql import make_identifier, SQL, IDENT_RE
def _sql_table(table) -> SQL:
def _sql_table(table: str | SQL | None) -> SQL | None:
""" Wrap an optional table as an SQL object. """
if isinstance(table, str):
return SQL.identifier(table) if IDENT_RE.match(table) else SQL(f"({table})")
return table
def _sql_from_table(alias, table) -> SQL:
def _sql_from_table(alias: str, table: SQL | None) -> SQL:
""" Return a FROM clause element from ``alias`` and ``table``. """
if table is None:
return SQL.identifier(alias)
return SQL("%s AS %s", table, SQL.identifier(alias))
def _sql_from_join(kind, alias, table, condition) -> SQL:
def _sql_from_join(kind: SQL, alias: str, table: SQL | None, condition: SQL) -> SQL:
""" Return a FROM clause element for a JOIN. """
return SQL("%s %s ON (%s)", kind, _sql_from_table(alias, table), condition)
@@ -61,7 +59,7 @@ class Query(object):
:param table: a table expression (``str`` or ``SQL`` object), optional
"""
def __init__(self, cr, alias: str, table: Union[str, SQL, None] = None):
def __init__(self, cr, alias: str, table: (str | SQL | None) = None):
# database cursor
self._cr = cr
@@ -86,13 +84,13 @@ class Query(object):
""" Return an alias based on ``alias`` and ``link``. """
return _generate_table_alias(alias, link)
def add_table(self, alias: str, table: Union[str, SQL, None] = None):
def add_table(self, alias: str, table: (str | SQL | None) = None):
""" Add a table with a given alias to the from clause. """
assert alias not in self._tables and alias not in self._joins, f"Alias {alias!r} already in {self}"
self._tables[alias] = _sql_table(table)
self._ids = None
def add_join(self, kind: str, alias: str, table: Union[str, SQL, None], condition: SQL):
def add_join(self, kind: str, alias: str, table: str | SQL | None, condition: SQL):
""" Add a join clause with the given alias, table and condition. """
sql_kind = _SQL_JOINS.get(kind.upper())
assert sql_kind is not None, f"Invalid JOIN type {kind!r}"
@@ -105,7 +103,7 @@ class Query(object):
self._joins[alias] = (sql_kind, table, condition)
self._ids = None
def add_where(self, where_clause: Union[str, SQL], where_params=()):
def add_where(self, where_clause: str | SQL, where_params=()):
""" Add a condition to the where clause. """
self._where_clauses.append(SQL(where_clause, *where_params))
self._ids = None
@@ -170,7 +168,7 @@ class Query(object):
""" Return whether the query is known to return nothing. """
return self._ids == ()
def select(self, *args: Union[str, SQL]) -> SQL:
def select(self, *args: str | SQL) -> SQL:
""" Return the SELECT query as an ``SQL`` object. """
sql_args = map(SQL, args) if args else [SQL.identifier(self.table, 'id')]
return SQL(
@@ -183,7 +181,7 @@ class Query(object):
SQL(" OFFSET %s", self.offset) if self.offset else SQL(),
)
def subselect(self, *args: Union[str, SQL]) -> SQL:
def subselect(self, *args: str | SQL) -> SQL:
""" Similar to :meth:`.select`, but for sub-queries.
This one avoids the ORDER BY clause when possible,
and includes parentheses around the subquery.
+5 -4
View File
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Part of Odoo. See LICENSE file for full copyright and licensing details.
# pylint: disable=sql-injection
from __future__ import annotations
import enum
import json
@@ -8,7 +9,7 @@ import logging
import re
from binascii import crc32
from collections import defaultdict
from typing import Iterable, Optional, Union
from typing import Iterable, Union
import psycopg2
@@ -56,7 +57,7 @@ class SQL:
__slots__ = ('__code', '__args')
# pylint: disable=keyword-arg-before-vararg
def __new__(cls, code: Union[str, "SQL"] = "", /, *args):
def __new__(cls, code: (str | SQL) = "", /, *args):
if isinstance(code, SQL):
return code
# validate the format of code
@@ -105,7 +106,7 @@ class SQL:
yield self.code
yield self.params
def join(self, args: Iterable) -> "SQL":
def join(self, args: Iterable) -> SQL:
""" Join SQL objects or parameters with ``self`` as a separator. """
args = list(args)
# optimizations for special cases
@@ -122,7 +123,7 @@ class SQL:
return SQL("%s" * len(items), *items)
@classmethod
def identifier(cls, name: str, subname: Optional[str] = None) -> "SQL":
def identifier(cls, name: str, subname: (str | None) = None) -> SQL:
""" Return an SQL object that represents an identifier. """
assert IDENT_RE.match(name), f"{name!r} invalid for SQL.identifier()"
if subname is None: