"""Reusable mixins for INSERT/UPDATE/DELETE builders."""
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, cast
from mypy_extensions import trait
from sqlglot import exp
from typing_extensions import Self
from sqlspec.builder._base import BuiltQuery, QueryBuilder
from sqlspec.builder._parsing_utils import (
_mark_explicit_table_quotes,
extract_expression,
extract_sql_object_expression,
)
from sqlspec.exceptions import SQLBuilderError
from sqlspec.protocols import SQLBuilderProtocol
from sqlspec.utils.serializers import schema_dump
from sqlspec.utils.type_guards import has_expression_and_sql, has_parameter_builder, is_dict
if TYPE_CHECKING:
from sqlspec.builder._column import Column
from sqlspec.builder._expression_wrappers import ExpressionWrapper
from sqlspec.builder._select import Case
__all__ = (
"DeleteFromClauseMixin",
"InsertFromSelectMixin",
"InsertIntoClauseMixin",
"InsertValuesMixin",
"ReturningClauseMixin",
"UpdateFromClauseMixin",
"UpdateSetClauseMixin",
"UpdateTableClauseMixin",
)
ARG_PAIR_COUNT = 2
SINGLE_VALUE_COUNT = 1
[docs]
@trait
class DeleteFromClauseMixin:
"""Mixin providing FROM clause support for DELETE builders."""
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
def from_(self, table: str) -> Self:
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Delete())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Delete):
msg = f"Base expression for Delete is {type(current_expr).__name__}, expected Delete."
raise SQLBuilderError(msg)
assert current_expr is not None
current_expr.set("this", _mark_explicit_table_quotes(exp.to_table(table), table))
return self
[docs]
@trait
class InsertIntoClauseMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
def into(self, table: str) -> Self:
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Insert())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Insert):
msg = "Cannot set target table on a non-INSERT expression."
raise SQLBuilderError(msg)
assert current_expr is not None
current_expr.set("this", _mark_explicit_table_quotes(exp.to_table(table), table))
return self
[docs]
@trait
class InsertValuesMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
_columns: list[str]
def columns(self, *columns: str | exp.Expr) -> Self:
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Insert())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Insert):
msg = "Cannot set columns on a non-INSERT expression."
raise SQLBuilderError(msg)
assert current_expr is not None
current_this = current_expr.args.get("this")
if current_this is None:
msg = "Table must be set using .into() before setting columns."
raise SQLBuilderError(msg)
if columns:
identifiers = [exp.to_identifier(col) if isinstance(col, str) else col for col in columns]
table_expression = current_this.this if isinstance(current_this, exp.Schema) else current_this
current_expr.set("this", exp.Schema(this=table_expression.copy(), expressions=identifiers))
elif isinstance(current_this, exp.Schema):
table_name = current_this.this
current_expr.set("this", exp.Table(this=table_name))
try:
cols = self._columns
if not columns:
cols.clear()
else:
cols[:] = [col if isinstance(col, str) else str(col) for col in columns]
except AttributeError:
pass
return self
def values(self, *values: Any, **kwargs: Any) -> Self:
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Insert())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Insert):
msg = "Cannot add values to a non-INSERT expression."
raise SQLBuilderError(msg)
assert current_expr is not None
if current_expr.args.get("this") is None:
msg = "The target table must be set using .into() before adding values."
raise SQLBuilderError(msg)
builder = cast("SQLBuilderProtocol", self)
positional_values = list(values)
if len(positional_values) == SINGLE_VALUE_COUNT and is_dict(positional_values[0]) and not kwargs:
kwargs = positional_values[0]
positional_values = []
if kwargs and positional_values:
msg = "Cannot mix positional values with keyword values."
raise SQLBuilderError(msg)
row_expressions: list[exp.Expr] = []
column_defs: list[str] = list(self._columns or [])
if kwargs:
if not column_defs:
self.columns(*kwargs.keys())
column_defs = list(self._columns or [])
for col, val in kwargs.items():
if isinstance(val, exp.Expr):
row_expressions.append(val)
continue
if has_expression_and_sql(val):
row_expressions.append(extract_sql_object_expression(val, builder=self))
continue
column_name = str(col).split(".")[-1]
placeholder, _ = builder.create_placeholder(val, column_name)
row_expressions.append(placeholder)
else:
if column_defs and len(positional_values) != len(column_defs):
msg = (
f"Number of values ({len(positional_values)}) does not match the number of specified columns "
f"({len(column_defs)})."
)
raise SQLBuilderError(msg)
for index, raw_value in enumerate(positional_values):
if isinstance(raw_value, exp.Expr):
row_expressions.append(raw_value)
elif has_expression_and_sql(raw_value):
row_expressions.append(extract_sql_object_expression(raw_value, builder=self))
else:
if column_defs and index < len(column_defs):
column_token = column_defs[index]
column_name = column_token.rsplit(".", maxsplit=1)[-1]
else:
column_name = f"value_{index + 1}"
placeholder, _ = builder.create_placeholder(raw_value, column_name)
row_expressions.append(placeholder)
values_node = current_expr.args.get("expression")
tuple_expression = exp.Tuple(expressions=row_expressions)
if isinstance(values_node, exp.Values):
values_node.append("expressions", tuple_expression)
else:
current_expr.set("expression", exp.Values(expressions=[tuple_expression]))
return self
def add_values(self, values: Sequence[Any]) -> Self:
return self.values(*values)
[docs]
@trait
class InsertFromSelectMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
def from_select(self, select_builder: SQLBuilderProtocol) -> Self:
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Insert())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Insert):
msg = "Cannot set INSERT source on a non-INSERT expression."
raise SQLBuilderError(msg)
assert current_expr is not None
if current_expr.args.get("this") is None:
msg = "The target table must be set using .into() before adding values."
raise SQLBuilderError(msg)
subquery_parameters = select_builder.parameters
if subquery_parameters:
builder_with_params = cast("SQLBuilderProtocol", self)
for param_name, param_value in subquery_parameters.items():
builder_with_params.add_parameter(param_value, name=param_name)
select_expr = select_builder.get_expression()
if select_expr and isinstance(select_expr, exp.Select):
current_expr.set("expression", select_expr.copy())
else:
msg = "SelectBuilder must have a valid SELECT expression."
raise SQLBuilderError(msg)
return self
[docs]
@trait
class UpdateTableClauseMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
def table(self, table_name: str, alias: str | None = None) -> Self:
current_expr = self.get_expression()
if current_expr is None or not isinstance(current_expr, exp.Update):
self.set_expression(exp.Update(this=None, expressions=[]))
current_expr = self.get_expression()
assert current_expr is not None
table_expr: exp.Expr = _mark_explicit_table_quotes(exp.to_table(table_name, alias=alias), table_name)
current_expr.set("this", table_expr)
return self
[docs]
@trait
class UpdateSetClauseMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
def _process_update_value(self, val: Any, col: Any) -> exp.Expr:
if isinstance(val, exp.Expr):
return val
if has_parameter_builder(val):
subquery = val.build()
sql_text = subquery.sql if isinstance(subquery, BuiltQuery) else str(subquery)
query_builder = cast("QueryBuilder", self)
value_expr = exp.paren(exp.maybe_parse(sql_text, dialect=query_builder.dialect))
for p_name, p_value in val.parameters.items():
query_builder.add_parameter(p_value, name=p_name)
return value_expr
if has_expression_and_sql(val):
return extract_sql_object_expression(val, builder=self)
sql_builder = cast("SQLBuilderProtocol", self)
column_name = col if isinstance(col, str) else str(col)
if "." in column_name:
column_name = column_name.split(".")[-1]
placeholder, _ = sql_builder.create_placeholder(val, column_name)
return placeholder
def set(self, *args: Any, **kwargs: Any) -> Self:
if not args and not kwargs:
return self
current_expr = self.get_expression()
if current_expr is None:
self.set_expression(exp.Update())
current_expr = self.get_expression()
if not isinstance(current_expr, exp.Update):
msg = "Cannot add SET clause to non-UPDATE expression."
raise SQLBuilderError(msg)
assert current_expr is not None
assignments: list[exp.Expr] = []
if len(args) == ARG_PAIR_COUNT and not kwargs:
col, val = args
col_expr = col if isinstance(col, exp.Column) else exp.column(col)
assignments.append(exp.EQ(this=col_expr, expression=self._process_update_value(val, col)))
elif (len(args) == SINGLE_VALUE_COUNT and isinstance(args[0], Mapping)) or kwargs:
all_values = dict(args[0] if args else {}, **kwargs)
for col, val in all_values.items():
assignments.append(exp.EQ(this=exp.column(col), expression=self._process_update_value(val, col)))
else:
msg = "Invalid arguments for set(): use (column, value), mapping, or kwargs."
raise SQLBuilderError(msg)
existing = current_expr.args.get("expressions", [])
current_expr.set("expressions", existing + assignments)
return self
[docs]
def set_from(self, data: Any, *, exclude_unset: bool = True) -> Self:
"""Set columns from a dict, dataclass, msgspec.Struct, Pydantic model, or attrs class.
Schema instances are normalised via :func:`sqlspec.utils.serializers.schema_dump`
with ``wire_format=False`` and dispatched to :meth:`set` via keyword unpack. The
dict shape uses Python attribute names regardless of msgspec ``rename=`` or
Pydantic ``Field(alias=...)``.
Args:
data: A dict, dataclass instance, ``msgspec.Struct``, ``pydantic.BaseModel``,
or ``attrs``-decorated class instance.
exclude_unset: If True, exclude fields that were never set (msgspec UNSET,
Pydantic ``model_fields_set``, dataclass empty defaults). No-op for attrs.
Returns:
The current builder instance for method chaining.
"""
payload = schema_dump(data, exclude_unset=exclude_unset, wire_format=False)
return self.set(**payload)
[docs]
@trait
class UpdateFromClauseMixin:
__slots__ = ()
def get_expression(self) -> exp.Expr | None: ...
def set_expression(self, expression: exp.Expr) -> None: ...
[docs]
def from_(self, table: str | exp.Expr | Any, alias: str | None = None) -> Self:
"""Add a table or subquery to the UPDATE statement's FROM clause.
Args:
table: Target table name, expression, or builder instance.
alias: Optional alias for the source table or subquery.
Returns:
The current builder instance for method chaining.
Raises:
SQLBuilderError: If called on a non-UPDATE expression or with an unsupported table type.
"""
current_expr = self.get_expression()
if current_expr is None or not isinstance(current_expr, exp.Update):
msg = "Cannot add FROM clause to non-UPDATE expression. Set the main table first."
raise SQLBuilderError(msg)
table_expr: exp.Expr
if isinstance(table, str):
table_expr = exp.to_table(table, alias=alias)
elif isinstance(table, exp.Expr):
if isinstance(table, (exp.Select, exp.SetOperation)):
table_expr = exp.Subquery(this=table.copy())
if alias:
table_expr = exp.alias_(table_expr, alias, table=True)
else:
table_expr = exp.alias_(table.copy(), alias, table=True) if alias else table.copy()
elif (
hasattr(table, "build")
or hasattr(table, "to_statement")
or hasattr(table, "get_expression")
or hasattr(table, "_expression")
):
raw_expression = None
if hasattr(table, "_build_final_expression"):
raw_expression = table._build_final_expression(copy=True)
elif hasattr(table, "get_expression"):
raw_expression = table.get_expression()
elif hasattr(table, "_expression"):
raw_expression = table._expression
if raw_expression is None:
msg = "Subquery builder has no expression to include in FROM clause."
raise SQLBuilderError(msg)
subquery_copy = (
raw_expression
if isinstance(table, QueryBuilder)
else raw_expression.copy()
if hasattr(raw_expression, "copy")
else raw_expression
)
base_builder = cast("QueryBuilder", self)
builder_alias = getattr(table, "alias_name", None) or getattr(table, "alias", None)
if not isinstance(builder_alias, str):
builder_alias = None
if not builder_alias and hasattr(raw_expression, "alias_or_name"):
builder_alias = raw_expression.alias_or_name
effective_alias = alias or builder_alias or "subquery"
subquery_params = getattr(table, "parameters", {})
if subquery_params and isinstance(subquery_params, dict):
param_mapping = base_builder._merge_cte_parameters(effective_alias, subquery_params)
if param_mapping:
subquery_copy = base_builder._update_placeholders(subquery_copy, param_mapping)
if isinstance(subquery_copy, exp.Values):
if alias:
cols: list[str] = []
existing_alias = subquery_copy.args.get("alias")
source_columns = getattr(table, "columns", None)
if existing_alias and existing_alias.args.get("columns"):
cols = [c.name for c in existing_alias.args["columns"]]
elif isinstance(source_columns, (list, tuple)):
cols = [str(c) for c in source_columns]
table_expr = exp.alias_(subquery_copy, alias, table=cols or False)
else:
table_expr = subquery_copy
elif isinstance(subquery_copy, exp.Subquery):
table_expr = exp.alias_(subquery_copy, alias, table=True) if alias else subquery_copy
elif isinstance(subquery_copy, (exp.Select, exp.SetOperation)):
table_expr = exp.Subquery(this=subquery_copy)
if alias or builder_alias:
table_expr = exp.alias_(table_expr, alias or builder_alias, table=True)
else:
msg = "UPDATE FROM builder sources must be SELECT, VALUES, or subquery expressions."
raise SQLBuilderError(msg)
else:
msg = f"Unsupported table type for FROM clause: {type(table)}"
raise SQLBuilderError(msg)
from_clause = current_expr.args.get("from_")
if from_clause is None:
current_expr.set("from_", exp.From(this=table_expr))
else:
from_table = from_clause.this
from_table.append("joins", exp.Join(this=table_expr))
return self
[docs]
@trait
class ReturningClauseMixin:
"""Mixin providing RETURNING clause support for DML builders."""
__slots__ = ()
_expression: exp.Expr | None
[docs]
def returning(self, *columns: "str | exp.Expr | Column | ExpressionWrapper | Case") -> Self:
"""Add RETURNING clause to the DML statement.
Args:
*columns: Columns or expressions to return.
Returns:
The builder instance for method chaining.
Raises:
SQLBuilderError: If expression not initialized or not DML.
"""
if self._expression is None:
msg = "Cannot add RETURNING: expression not initialized."
raise SQLBuilderError(msg)
if not isinstance(self._expression, (exp.Insert, exp.Update, exp.Delete)):
msg = "RETURNING only supported for INSERT, UPDATE, DELETE statements."
raise SQLBuilderError(msg)
returning_exprs = [extract_expression(col) for col in columns]
self._expression.set("returning", exp.Returning(expressions=returning_exprs))
return self