From f07e050c9ce4afdeb9c0c136dbcc547f7e5ac7b8 Mon Sep 17 00:00:00 2001 From: Mike Bayer Date: Mon, 29 Apr 2019 23:26:36 -0400 Subject: Implement new ClauseElement role and coercion system A major refactoring of all the functions handle all detection of Core argument types as well as perform coercions into a new class hierarchy based on "roles", each of which identify a syntactical location within a SQL statement. In contrast to the ClauseElement hierarchy that identifies "what" each object is syntactically, the SQLRole hierarchy identifies the "where does it go" of each object syntactically. From this we define a consistent type checking and coercion system that establishes well defined behviors. This is a breakout of the patch that is reorganizing select() constructs to no longer be in the FromClause hierarchy. Also includes a rename of as_scalar() into scalar_subquery(); deprecates automatic coercion to scalar_subquery(). Partially-fixes: #4617 Change-Id: I26f1e78898693c6b99ef7ea2f4e7dfd0e8e1a1bd --- lib/sqlalchemy/dialects/mssql/base.py | 5 ++--- lib/sqlalchemy/dialects/mysql/base.py | 6 ++++-- lib/sqlalchemy/dialects/postgresql/base.py | 7 ++++-- lib/sqlalchemy/dialects/postgresql/ext.py | 34 +++++++++++++++++++++++------- 4 files changed, 37 insertions(+), 15 deletions(-) (limited to 'lib/sqlalchemy/dialects') diff --git a/lib/sqlalchemy/dialects/mssql/base.py b/lib/sqlalchemy/dialects/mssql/base.py index d2c84f446..00a110aa2 100644 --- a/lib/sqlalchemy/dialects/mssql/base.py +++ b/lib/sqlalchemy/dialects/mssql/base.py @@ -672,6 +672,7 @@ from ... import util from ...engine import default from ...engine import reflection from ...sql import compiler +from ...sql import elements from ...sql import expression from ...sql import quoted_name from ...sql import util as sql_util @@ -1671,9 +1672,7 @@ class MSSQLCompiler(compiler.SQLCompiler): # translate for schema-qualified table aliases t = self._schema_aliased_table(column.table) if t is not None: - converted = expression._corresponding_column_or_error( - t, column - ) + converted = elements._corresponding_column_or_error(t, column) if add_to_result_map is not None: add_to_result_map( column.name, diff --git a/lib/sqlalchemy/dialects/mysql/base.py b/lib/sqlalchemy/dialects/mysql/base.py index 44f90c47c..9cae3c689 100644 --- a/lib/sqlalchemy/dialects/mysql/base.py +++ b/lib/sqlalchemy/dialects/mysql/base.py @@ -788,8 +788,10 @@ from ... import types as sqltypes from ... import util from ...engine import default from ...engine import reflection +from ...sql import coercions from ...sql import compiler from ...sql import elements +from ...sql import roles from ...types import BINARY from ...types import BLOB from ...types import BOOLEAN @@ -1218,7 +1220,7 @@ class MySQLCompiler(compiler.SQLCompiler): def visit_on_duplicate_key_update(self, on_duplicate, **kw): if on_duplicate._parameter_ordering: parameter_ordering = [ - elements._column_as_key(key) + coercions.expect(roles.DMLColumnRole, key) for key in on_duplicate._parameter_ordering ] ordered_keys = set(parameter_ordering) @@ -1238,7 +1240,7 @@ class MySQLCompiler(compiler.SQLCompiler): val = on_duplicate.update.get(column.key) if val is None: continue - elif elements._is_literal(val): + elif coercions._is_literal(val): val = elements.BindParameter(None, val, type_=column.type) value_text = self.process(val.self_group(), use_schema=False) elif isinstance(val, elements.BindParameter) and val.type._isnull: diff --git a/lib/sqlalchemy/dialects/postgresql/base.py b/lib/sqlalchemy/dialects/postgresql/base.py index ceb624644..f18bec932 100644 --- a/lib/sqlalchemy/dialects/postgresql/base.py +++ b/lib/sqlalchemy/dialects/postgresql/base.py @@ -932,9 +932,11 @@ from ... import sql from ... import util from ...engine import default from ...engine import reflection +from ...sql import coercions from ...sql import compiler from ...sql import elements from ...sql import expression +from ...sql import roles from ...sql import sqltypes from ...sql import util as sql_util from ...types import BIGINT @@ -1774,7 +1776,7 @@ class PGCompiler(compiler.SQLCompiler): col_key = c.key if col_key in set_parameters: value = set_parameters.pop(col_key) - if elements._is_literal(value): + if coercions._is_literal(value): value = elements.BindParameter(None, value, type_=c.type) else: @@ -1806,7 +1808,8 @@ class PGCompiler(compiler.SQLCompiler): else self.process(k, use_schema=False) ) value_text = self.process( - elements._literal_as_binds(v), use_schema=False + coercions.expect(roles.ExpressionElementRole, v), + use_schema=False, ) action_set_ops.append("%s = %s" % (key_text, value_text)) diff --git a/lib/sqlalchemy/dialects/postgresql/ext.py b/lib/sqlalchemy/dialects/postgresql/ext.py index 426028239..f9cbc945a 100644 --- a/lib/sqlalchemy/dialects/postgresql/ext.py +++ b/lib/sqlalchemy/dialects/postgresql/ext.py @@ -6,9 +6,12 @@ # the MIT License: http://www.opensource.org/licenses/mit-license.php from .array import ARRAY +from ... import util +from ...sql import coercions from ...sql import elements from ...sql import expression from ...sql import functions +from ...sql import roles from ...sql.schema import ColumnCollectionConstraint @@ -50,16 +53,18 @@ class aggregate_order_by(expression.ColumnElement): __visit_name__ = "aggregate_order_by" def __init__(self, target, *order_by): - self.target = elements._literal_as_binds(target) + self.target = coercions.expect(roles.ExpressionElementRole, target) _lob = len(order_by) if _lob == 0: raise TypeError("at least one ORDER BY element is required") elif _lob == 1: - self.order_by = elements._literal_as_binds(order_by[0]) + self.order_by = coercions.expect( + roles.ExpressionElementRole, order_by[0] + ) else: self.order_by = elements.ClauseList( - *order_by, _literal_as_text=elements._literal_as_binds + *order_by, _literal_as_text_role=roles.ExpressionElementRole ) def self_group(self, against=None): @@ -166,7 +171,10 @@ class ExcludeConstraint(ColumnCollectionConstraint): expressions, operators = zip(*elements) for (expr, column, strname, add_element), operator in zip( - self._extract_col_expression_collection(expressions), operators + coercions.expect_col_expression_collection( + roles.DDLConstraintColumnRole, expressions + ), + operators, ): if add_element is not None: columns.append(add_element) @@ -177,8 +185,6 @@ class ExcludeConstraint(ColumnCollectionConstraint): # backwards compat self.operators[name] = operator - expr = expression._literal_as_column(expr) - render_exprs.append((expr, name, operator)) self._render_exprs = render_exprs @@ -193,9 +199,21 @@ class ExcludeConstraint(ColumnCollectionConstraint): self.using = kw.get("using", "gist") where = kw.get("where") if where is not None: - self.where = expression._literal_as_text( - where, allow_coercion_to_text=True + self.where = coercions.expect(roles.StatementOptionRole, where) + + def _set_parent(self, table): + super(ExcludeConstraint, self)._set_parent(table) + + self._render_exprs = [ + ( + expr if isinstance(expr, elements.ClauseElement) else colexpr, + name, + operator, + ) + for (expr, name, operator), colexpr in util.zip_longest( + self._render_exprs, self.columns ) + ] def copy(self, **kw): elements = [(col, self.operators[col]) for col in self.columns.keys()] -- cgit v1.2.1