[IMP] core: type annotations in SQL wrapper and Query
Part-of: odoo/odoo#138019
This commit is contained in:
+9
-11
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user