Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
17 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions docs/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -212,6 +212,9 @@ v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs param
consumption without intermediate tuple relays.
(`#771 <https://github.com/litestar-org/sqlspec/pull/771>`_)

* ``sql.values`` creates a :class:`~sqlspec.builder.Values` builder for parameterized bulk row lists rather than resolving as a column named ``values``. Use ``sql.column("values")`` to construct column expressions referencing that identifier.
(`#779 <https://github.com/litestar-org/sqlspec/pull/779>`_)

**Fixed:**

* Preserve JSON objects and arrays as individual query parameters after placeholder conversion,
Expand All @@ -233,6 +236,12 @@ v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs param
deduplicate concurrent event IDs with a single key-range-locked statement.
(`#781 <https://github.com/litestar-org/sqlspec/pull/781>`_)

* ``Update.from_()`` accepts query builders and subqueries with parameter merging instead of raising a runtime type error, and dialect checks raise a descriptive ``SQLBuilderError`` on dialects without native ``FROM`` clause support.
(`#779 <https://github.com/litestar-org/sqlspec/pull/779>`_)

* CTEs registered via ``with_cte()`` or ``with_()`` render on ``Update`` and ``Delete`` statements and merge bound parameters without inlining or dropping explicit common table expressions.
(`#779 <https://github.com/litestar-org/sqlspec/pull/779>`_)

* Query builder keeps ``ON CONFLICT ... DO UPDATE`` and ``ON DUPLICATE KEY UPDATE`` assignments in written
order. Assignments such as ``do_update(name=exp.column("name", table="excluded"))`` no longer render
reversed. Conflict targets and update columns are quoted, allowing reserved words (e.g. ``order``, ``group``)
Expand Down
7 changes: 7 additions & 0 deletions docs/reference/builder/queries.rst
Original file line number Diff line number Diff line change
Expand Up @@ -158,3 +158,10 @@ Joins
.. autoclass:: JoinBuilder
:members:
:show-inheritance:

Values
======

.. autoclass:: Values
:members:
:show-inheritance:
62 changes: 62 additions & 0 deletions docs/usage/query_builder.rst
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,68 @@ Inserts and Updates
:dedent: 4
:no-upgrade:

UPDATE ... FROM and CTEs
------------------------

Builder queries support ``UPDATE ... FROM`` with subqueries or Common Table Expressions (CTEs),
enabling queue claim statements and batch updates across supported dialects.

.. code-block:: python

from sqlspec import sql

claim_query = (
sql.update("tasks")
.set(status="processing")
.from_(
sql.select("id")
.from_("tasks")
.where_eq("status", "pending")
.limit(1)
.for_update(skip_locked=True),
alias="sub",
)
.where("tasks.id = sub.id")
.returning(sql.column("id", table="tasks"))
)

Dialect Support Matrix
~~~~~~~~~~~~~~~~~~~~~~

.. list-table::
:header-rows: 1

* - Dialect
- UPDATE ... FROM Support
- Notes
* - PostgreSQL
- Yes
- Native ``UPDATE ... FROM`` with ``RETURNING`` and row locking
* - CockroachDB
- Yes
- Native ``UPDATE ... FROM``
* - SQLite
- Yes
- Native ``UPDATE ... FROM`` (SQLite 3.33.0+)
* - DuckDB
- Yes
- Native ``UPDATE ... FROM``
* - SQL Server (MSSQL)
- Yes
- Native ``UPDATE ... FROM``
* - MySQL / MariaDB
- No
- Raises ``SQLBuilderError``; use multi-table join update or MERGE
* - Oracle
- No
- Builder conservatively raises ``SQLBuilderError``; use MERGE or raw SQL on Oracle 23+
* - Spanner
- No
- Raises ``SQLBuilderError``
* - BigQuery
- Yes
- Native ``UPDATE ... FROM``; a ``WHERE`` condition is required

Upserts (ON CONFLICT)
---------------------

Expand Down
2 changes: 2 additions & 0 deletions sqlspec/builder/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,7 @@
)
from sqlspec.builder._temporal import create_temporal_table, register_version_generators
from sqlspec.builder._update import Update
from sqlspec.builder._values import Values
from sqlspec.builder._vector_distance import VectorDistance
from sqlspec.exceptions import SQLBuilderError

Expand Down Expand Up @@ -151,6 +152,7 @@
"UpdateFromClauseMixin",
"UpdateSetClauseMixin",
"UpdateTableClauseMixin",
"Values",
"VectorDistance",
"WhereClauseMixin",
"WindowFunctionBuilder",
Expand Down
161 changes: 136 additions & 25 deletions sqlspec/builder/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@
from sqlglot.dialects.dialect import DialectType
from sqlglot.errors import ParseError as SQLGlotParseError
from sqlglot.optimizer import RULES, optimize
from sqlglot.optimizer.eliminate_ctes import eliminate_ctes as _eliminate_ctes_rule
from sqlglot.optimizer.merge_subqueries import merge_subqueries as _merge_subqueries_rule
from sqlglot.optimizer.normalize_identifiers import normalize_identifiers as _normalize_identifiers_rule
from sqlglot.optimizer.optimize_joins import optimize_joins as _optimize_joins_rule
from sqlglot.optimizer.pushdown_predicates import pushdown_predicates as _pushdown_predicates_rule
Expand All @@ -38,7 +40,7 @@
from sqlspec.data_dictionary import get_dialect_config
from sqlspec.exceptions import SQLBuilderError
from sqlspec.utils.logging import get_logger
from sqlspec.utils.type_guards import has_expression_and_parameters, has_name, has_with_method, is_expression
from sqlspec.utils.type_guards import has_expression_and_parameters, has_name, is_expression
from sqlspec.utils.uuids import uuid4

__all__ = ("BuiltQuery", "ExpressionBuilder", "QueryBuilder")
Expand Down Expand Up @@ -316,14 +318,16 @@ def _build_final_expression(self, *, copy: bool = False) -> exp.Expr:
return base_expression

final_expression: exp.Expr = base_expression
if has_with_method(final_expression):
for alias, cte_node in self._with_ctes.items():
final_expression = cast("Any", final_expression).with_(alias, as_=cte_node.args["this"], copy=False)
return cast("exp.Expr", final_expression)

if "with_" in type(final_expression).arg_types:
existing_with = final_expression.args.get("with_")
if existing_with is None:
final_expression.set("with_", exp.With(expressions=list(self._with_ctes.values())))
else:
for cte_node in self._with_ctes.values():
if cte_node not in existing_with.expressions:
existing_with.append("expressions", cte_node)

if any(cte.meta.get("recursive") for cte in self._with_ctes.values()):
final_expression.args["with_"].set("recursive", True)
return final_expression

def _spawn_like_self(self: Self) -> Self:
Expand All @@ -337,35 +341,66 @@ def _spawn_like_self(self: Self) -> Self:
simplify_expressions=self.simplify_expressions,
)

def _resolve_cte_query(self, alias: str, query: "QueryBuilder | exp.Select | str") -> exp.Select:
"""Resolve a CTE query into a Select expression with merged parameters."""
def _resolve_cte_query(
self, alias: str, query: "QueryBuilder | exp.Select | exp.SetOperation | exp.Values | str | Any"
) -> exp.Expr:
"""Resolve a CTE query into a Select or Values expression with merged parameters."""
if isinstance(query, QueryBuilder):
query_expr = query.get_expression()
query_expr = query._build_final_expression(copy=True)
if query_expr is None:
self._raise_cte_query_error(alias, "query builder has no expression")
if not isinstance(query_expr, exp.Select):
self._raise_cte_query_error(alias, f"expression must be a Select, got {type(query_expr).__name__}")
cte_select_expression = query_expr.copy()
if not isinstance(query_expr, (exp.Select, exp.SetOperation, exp.Values)):
self._raise_cte_query_error(
alias, f"expression must be a Select or Values, got {type(query_expr).__name__}"
)
cte_select_expression: exp.Expr = query_expr.copy()
if isinstance(cte_select_expression, exp.Values) and cte_select_expression.args.get("alias"):
cte_select_expression = cte_select_expression.copy()
cte_select_expression.set("alias", None)
param_mapping = self._merge_cte_parameters(alias, query.parameters)
updated_expression = self._update_placeholders(cte_select_expression, param_mapping)
if not isinstance(updated_expression, exp.Select): # pragma: no cover
msg = "CTE placeholder update produced non-select expression"
raise SQLBuilderError(msg)
return updated_expression
if param_mapping:
cte_select_expression = self._update_placeholders(cte_select_expression, param_mapping)
return cte_select_expression

if hasattr(query, "get_expression"):
raw_query = cast("Any", query)
raw_query_expr = (
raw_query._build_final_expression(copy=True)
if hasattr(raw_query, "_build_final_expression")
else raw_query.get_expression()
)
if raw_query_expr is None:
self._raise_cte_query_error(alias, "query builder has no expression")
if not isinstance(raw_query_expr, (exp.Select, exp.SetOperation, exp.Values)):
self._raise_cte_query_error(
alias, f"expression must be a Select or Values, got {type(raw_query_expr).__name__}"
)
cte_duck_expression: exp.Expr = raw_query_expr.copy()
if isinstance(cte_duck_expression, exp.Values) and cte_duck_expression.args.get("alias"):
cte_duck_expression = cte_duck_expression.copy()
cte_duck_expression.set("alias", None)
if hasattr(raw_query, "parameters"):
param_mapping = self._merge_cte_parameters(alias, raw_query.parameters)
if param_mapping:
cte_duck_expression = self._update_placeholders(cte_duck_expression, param_mapping)
return cte_duck_expression

if isinstance(query, str):
try:
parsed_expression = sqlglot.parse_one(query, read=self.dialect_name)
except SQLGlotParseError as e: # pragma: no cover
self._raise_cte_parse_error(e)
if not isinstance(parsed_expression, exp.Select):
if not isinstance(parsed_expression, (exp.Select, exp.SetOperation, exp.Values)):
self._raise_cte_query_error(
alias, f"query string must parse to SELECT, got {type(parsed_expression).__name__}"
alias, f"query string must parse to SELECT or VALUES, got {type(parsed_expression).__name__}"
)
return parsed_expression

if isinstance(query, exp.Select):
return query
if isinstance(query, (exp.Select, exp.SetOperation, exp.Values)):
cte_expression = query.copy()
if isinstance(cte_expression, exp.Values):
cte_expression.set("alias", None)
return cte_expression

self._raise_cte_query_error(alias, f"invalid query type: {type(query).__name__}")
msg = "Unreachable"
Expand Down Expand Up @@ -579,13 +614,21 @@ def _cache_key(self, config: "StatementConfig | None" = None) -> str:
)
return f"builder:{fingerprint}"

def with_cte(self: Self, alias: str, query: "QueryBuilder | exp.Select | str") -> Self:
def with_cte(
self: Self,
alias: str,
query: "QueryBuilder | exp.Select | exp.SetOperation | exp.Values | str | Any",
recursive: bool = False,
columns: "list[str] | None" = None,
) -> Self:
"""Adds a Common Table Expression (CTE) to the query.

Args:
alias: The alias for the CTE.
query: The CTE query, which can be another QueryBuilder instance,
a raw SQL string, or a sqlglot Select expression.
a raw SQL string, or a sqlglot Select or Values expression.
recursive: Whether the CTE is recursive.
columns: Optional list of column aliases for the CTE.

Returns:
Self: The current builder instance for method chaining.
Expand All @@ -594,9 +637,49 @@ def with_cte(self: Self, alias: str, query: "QueryBuilder | exp.Select | str") -
self._raise_builder_error(f"CTE with alias '{alias}' already exists.")

cte_select_expression = self._resolve_cte_query(alias, query)
self._with_ctes[alias] = exp.CTE(this=cte_select_expression, alias=exp.to_table(alias))
cte_columns = columns
if cte_columns is None and isinstance(query, exp.Values):
values_alias = query.args.get("alias")
if isinstance(values_alias, exp.TableAlias):
cte_columns = [column.name for column in values_alias.columns]
if cte_columns is None:
query_cols = getattr(query, "columns", None)
query_private_cols = getattr(query, "_columns", None)
if isinstance(query_cols, (list, tuple)):
cte_columns = [str(c) for c in query_cols]
elif isinstance(query_private_cols, (list, tuple)):
cte_columns = [str(c) for c in query_private_cols]

if cte_columns:
alias_node: exp.Expr = exp.TableAlias(
this=exp.to_identifier(alias), columns=[exp.to_identifier(c) for c in cte_columns]
)
else:
alias_node = exp.to_table(alias)
self._with_ctes[alias] = exp.CTE(this=cte_select_expression, alias=alias_node)
self._with_ctes[alias].meta["recursive"] = recursive
return self

def with_(
self: Self,
name: str,
query: "QueryBuilder | exp.Select | exp.SetOperation | exp.Values | str | Any",
recursive: bool = False,
columns: "list[str] | None" = None,
) -> Self:
"""Alias for with_cte for parity across builders.

Args:
name: The alias/name for the CTE.
query: The CTE query expression or builder.
recursive: Whether the CTE is recursive.
columns: Optional list of column aliases for the CTE.

Returns:
Self: The current builder instance for method chaining.
"""
return self.with_cte(name, query, recursive=recursive, columns=columns)

def build(self, dialect: DialectType = None) -> "BuiltQuery":
"""Builds the SQL query string and parameters.

Expand All @@ -608,6 +691,7 @@ def build(self, dialect: DialectType = None) -> "BuiltQuery":
BuiltQuery: A dataclass containing the SQL string and parameters.
"""
final_expression = self._build_final_expression()
self._validate_update_from(final_expression, _resolve_dialect(dialect, self.dialect))

if self.enable_optimization and isinstance(final_expression, exp.Expr):
final_expression = self._optimize_expression(final_expression)
Expand Down Expand Up @@ -805,13 +889,20 @@ def _optimize_expression(self, expression: exp.Expr, *, force: bool = False) ->
if cached_optimized is not None:
return cast("exp.Expr", cached_optimized).copy()

# Qualification drops VALUES CTE column aliases without projecting replacements.
if any(isinstance(cte.this, exp.Values) and cte.alias_column_names for cte in expression.find_all(exp.CTE)):
return expression

excluded_rules = set()
if not self.optimize_joins:
excluded_rules.add(_optimize_joins_rule)
if not self.optimize_predicates:
excluded_rules.add(_pushdown_predicates_rule)
if not self.simplify_expressions:
excluded_rules.add(_simplify_rule)
if expression.args.get("with_") is not None or self._with_ctes:
excluded_rules.add(_eliminate_ctes_rule)
excluded_rules.add(_merge_subqueries_rule)

rules = RULES if not excluded_rules else tuple(rule for rule in RULES if rule not in excluded_rules)

Expand Down Expand Up @@ -914,10 +1005,30 @@ def _to_statement(self, config: "StatementConfig | None" = None) -> "SQL":
cache_entry = self._create_builder_cache_entry(config)
return self._statement_from_cache_entry(cache_entry, config)

def _validate_update_from(self, expression: exp.Expr, dialect: DialectType) -> None:
if not dialect or not any(node.args.get("from_") is not None for node in expression.find_all(exp.Update)):
return
from sqlspec.data_dictionary import get_dialect_config

dialect_name = (
dialect.lower() if isinstance(dialect, str) else type(Dialect.get_or_raise(dialect)).__name__.lower()
)
try:
config = get_dialect_config("spanner" if dialect_name == "spangres" else dialect_name)
except ValueError:
return
if not config.feature_flags.get("supports_update_from", True):
msg = (
f"Dialect '{dialect_name}' does not support UPDATE ... FROM clauses. "
"Consider using MERGE or a JOIN-based UPDATE instead."
)
raise SQLBuilderError(msg)

def _create_builder_cache_entry(self, config: "StatementConfig | None") -> "_BuilderCacheEntry":
dialect_override = config.dialect if config is not None else None
resolved_dialect = self._build_dialect(dialect_override)
statement_expression = self._build_final_expression(copy=True)
self._validate_update_from(statement_expression, resolved_dialect)

if self.enable_optimization and isinstance(statement_expression, exp.Expr):
statement_expression = self._optimize_expression(statement_expression)
Expand Down
Loading
Loading