diff --git a/docs/changelog.rst b/docs/changelog.rst index 7d267f9ee..e3e15a6c9 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -153,8 +153,20 @@ v0.63.0 - Transactions, table fixtures, SQL fragments, storage, and kwargs param **Changed:** +* Query builder validates dialect capabilities in ``build()`` and ``to_statement()`` instead of silently dropping unsupported clauses. + ``.for_update()`` and ``.for_share()`` raise ``SQLBuilderError`` on dialects without these locking clauses + (T-SQL, SQLite, DuckDB, BigQuery), and ``skip_locked=True`` validates against ``supports_skip_locked``. + ``.on_conflict()`` automatically transpiles to ``ON DUPLICATE KEY UPDATE`` for MySQL and MariaDB (with ``do_nothing()`` + rewriting to self-assignment), while raising ``SQLBuilderError`` suggesting ``sql.merge()`` on dialects lacking native + upsert support (Oracle, T-SQL, BigQuery). Spanner supports native upserts and plain ``FOR UPDATE``; + PostgreSQL-mode upserts validate assignment restrictions. Oracle rejects shared locks, while MariaDB renders + ``LOCK IN SHARE MODE``. Query builder ``build()`` also normalizes dialect aliases + (``mssql`` to ``tsql``, ``mariadb`` to ``mysql``, and ``cockroachdb`` to ``postgres``). + (`#778 `_) + * Decoupled the ``ParadeDB`` dialect so it inherits directly from ``Postgres`` rather than ``PGVector``, allowing clean independent combinations of vector search and BM25 extensions. + * Standardized PostgreSQL extension detection across all adapters on a single first-connection probe via :func:`~sqlspec.core.config_runtime.build_postgres_extension_probe_names`, removing ad-hoc ADK probe branches. diff --git a/docs/usage/query_builder.rst b/docs/usage/query_builder.rst index 0b4d8d7d8..d83ec0368 100644 --- a/docs/usage/query_builder.rst +++ b/docs/usage/query_builder.rst @@ -44,6 +44,16 @@ Upserts (ON CONFLICT) Use ``.on_conflict()`` to handle insert conflicts. Chain ``.do_nothing()`` to skip conflicting rows, or ``.do_update(**columns)`` to update them. +Dialects natively supporting ``ON CONFLICT`` (PostgreSQL, CockroachDB, SQLite, DuckDB, and Spanner) +render standard ``ON CONFLICT`` syntax. For MySQL and MariaDB, the builder automatically +transpiles ``.on_conflict().do_update()`` to ``ON DUPLICATE KEY UPDATE``, and ``.do_nothing()`` +to a no-op self-assignment (e.g., ``col = col``). This requires a conflict column or +explicit insert columns. MySQL handles conflicts on any unique key, regardless of the +requested conflict target; the no-op update can still fire update triggers. References +to ``excluded.column`` in update expressions become ``VALUES(column)``. Dialects without native upsert clauses +(Oracle, T-SQL / SQL Server, and BigQuery) raise :class:`~sqlspec.exceptions.SQLBuilderError` +in both ``build()`` and ``to_statement()`` advising the use of :func:`sql.merge`. + .. literalinclude:: /examples/builder/upsert.py :language: python :caption: ``upsert with on_conflict`` @@ -81,6 +91,20 @@ Joins Query Modifiers --------------- +Row-level locking clauses such as ``.for_update()`` and ``.for_share()`` are validated against +dialect capabilities at build time. On dialects without these locking clauses (T-SQL, SQLite, +DuckDB, and BigQuery), building a locked query raises :class:`~sqlspec.exceptions.SQLBuilderError`. +Oracle also rejects ``.for_share()``; MariaDB renders it as ``LOCK IN SHARE MODE`` +and rejects ``of=`` targets for all locking clauses. PostgreSQL key lock variants +are rejected on other dialect families. +Spanner supports plain ``FOR UPDATE`` in both SQL modes, but rejects shared locks, +``SKIP LOCKED``, ``NOWAIT``, and ``OF`` modifiers. Its PostgreSQL mode requires conflict +updates to assign every inserted column from the matching ``excluded`` column and +does not accept conflict predicates or named constraints. +Similarly, ``skip_locked=True`` is validated against the dialect's ``supports_skip_locked`` capability. +The builder also normalizes common dialect aliases during build (e.g., ``mssql`` to ``tsql``, +``mariadb`` to ``mysql``, and ``cockroachdb`` to ``postgres``). + .. literalinclude:: /examples/builder/query_modifiers.py :language: python :caption: ``where helpers + pagination`` diff --git a/sqlspec/builder/_base.py b/sqlspec/builder/_base.py index afa672d1a..7b3c10f1c 100644 --- a/sqlspec/builder/_base.py +++ b/sqlspec/builder/_base.py @@ -21,7 +21,7 @@ from typing_extensions import Self from sqlspec.builder._locking import register_lock_generator -from sqlspec.builder._parsing_utils import _resolve_dialect +from sqlspec.builder._parsing_utils import _normalize_dialect, _resolve_dialect from sqlspec.builder._vector_distance import has_vector_distance_ancestor from sqlspec.core import ( SQL, @@ -35,6 +35,7 @@ ) from sqlspec.core.filters import StatementFilter from sqlspec.core.hashing import _expression_cache_fingerprint +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 @@ -611,7 +612,8 @@ def build(self, dialect: DialectType = None) -> "BuiltQuery": if self.enable_optimization and isinstance(final_expression, exp.Expr): final_expression = self._optimize_expression(final_expression) - target_dialect = str(dialect) if dialect else self.dialect_name + target_dialect = self._build_dialect(dialect) + final_expression = self._prepare_dialect_expression(final_expression, target_dialect, dialect) try: if isinstance(final_expression, exp.Expr): @@ -631,9 +633,98 @@ def build(self, dialect: DialectType = None) -> "BuiltQuery": err_msg = f"Error generating SQL from expression: {e!s}" self._raise_builder_error(err_msg, e) - return BuiltQuery( - sql=sql_string, parameters=self._parameters.copy(), dialect=_resolve_dialect(dialect, self.dialect) - ) + return BuiltQuery(sql=sql_string, parameters=self._parameters.copy(), dialect=target_dialect) + + def _build_dialect(self, dialect: DialectType = None) -> str | None: + return _normalize_dialect(dialect or self.dialect) + + def _prepare_dialect_expression( + self, expression: exp.Expr, dialect: str | None, source_dialect: DialectType = None + ) -> exp.Expr: + """Validate and translate a copy so builds never mutate the builder AST.""" + if dialect is None: + return expression + try: + config = get_dialect_config("spanner" if dialect == "spangres" else dialect) + except ValueError: + return expression + if str(source_dialect or self.dialect).lower() == "mariadb" and expression.find(exp.Lock): + expression = expression.copy() + for lock in expression.find_all(exp.Lock): + if lock.expressions: + self._raise_builder_error("MariaDB locking clauses do not support OF targets.") + if not lock.args.get("update"): + lock.set("sqlspec_share_mode", True) + for lock in expression.find_all(exp.Lock): + if dialect in {"spanner", "spangres"} and ( + not lock.args.get("update") + or lock.args.get("wait") is not None + or lock.expressions + or lock.args.get("key") + ): + self._raise_builder_error(f"Dialect '{dialect}' supports only plain FOR UPDATE without lock modifiers.") + if lock.args.get("key") and dialect != "postgres": + self._raise_builder_error(f"Dialect '{dialect}' does not support PostgreSQL key lock modes.") + if dialect == "oracle" and not lock.args.get("update"): + self._raise_builder_error("Dialect 'oracle' does not support FOR SHARE.") + if dialect not in {"spanner", "spangres"} and config.get_feature_flag("supports_for_update") is False: + self._raise_builder_error(f"Dialect '{dialect}' does not support FOR UPDATE / row locking.") + if lock.args.get("wait") is False and config.get_feature_flag("supports_skip_locked") is False: + self._raise_builder_error(f"Dialect '{dialect}' does not support SKIP LOCKED.") + if dialect == "spangres": + self._validate_spangres_conflicts(expression) + if config.get_feature_flag("supports_on_conflict") is not False or not expression.find(exp.OnConflict): + return expression + if dialect != "mysql": + self._raise_builder_error(f"Dialect '{dialect}' does not support ON CONFLICT; use sql.merge() instead.") + expression = expression.copy() + for conflict in expression.find_all(exp.OnConflict): + if conflict.args.get("duplicate"): + continue + if any(conflict.args.get(key) for key in ("where", "index_predicate", "constraint")): + self._raise_builder_error("MySQL cannot preserve ON CONFLICT predicates or named constraints.") + assignments = conflict.args.get("expressions") + if str(conflict.args.get("action", "")).upper() == "DO NOTHING": + keys = conflict.args.get("conflict_keys") + insert = conflict.find_ancestor(exp.Insert) + schema = insert.this if insert is not None else None + columns = keys or (schema.expressions if isinstance(schema, exp.Schema) else None) + if not columns: + self._raise_builder_error("MySQL DO NOTHING requires a conflict column or explicit insert columns.") + column = exp.column(columns[0].name) + assignments = [exp.EQ(this=column, expression=column.copy())] + elif not assignments: + self._raise_builder_error("ON CONFLICT DO UPDATE requires at least one assignment.") + for assignment in assignments: + for column in list(assignment.find_all(exp.Column)): + if column.table.lower() == "excluded": + column.replace(exp.Anonymous(this="VALUES", expressions=[exp.column(column.name)])) + conflict.replace(exp.OnConflict(duplicate=True, action=exp.var("UPDATE"), expressions=assignments)) + return expression + + def _validate_spangres_conflicts(self, expression: exp.Expr) -> None: + """Check PostgreSQL-mode Spanner upsert restrictions visible in the AST.""" + for conflict in expression.find_all(exp.OnConflict): + if any(conflict.args.get(key) for key in ("where", "index_predicate", "constraint", "duplicate")): + self._raise_builder_error("Spanner PostgreSQL does not support these ON CONFLICT modifiers.") + assignments = conflict.args.get("expressions") or [] + if assignments: + insert = conflict.find_ancestor(exp.Insert) + schema = insert.this if insert is not None else None + assigned_columns = set() + for assignment in assignments: + value = assignment.expression + if ( + not isinstance(value, exp.Column) + or value.table.lower() != "excluded" + or value.name != assignment.this.name + ): + self._raise_builder_error("Spanner PostgreSQL conflict updates require excluded column values.") + assigned_columns.add(assignment.this.name) + if isinstance(schema, exp.Schema) and assigned_columns != { + column.name for column in schema.expressions + }: + self._raise_builder_error("Spanner PostgreSQL conflict updates must assign every inserted column.") def to_sql(self, show_parameters: bool = False, dialect: DialectType = None) -> str: """Return SQL string with optional parameter substitution. @@ -825,12 +916,16 @@ def _to_statement(self, config: "StatementConfig | None" = None) -> "SQL": def _create_builder_cache_entry(self, config: "StatementConfig | None") -> "_BuilderCacheEntry": dialect_override = config.dialect if config is not None else None - resolved_dialect = _resolve_dialect(dialect_override, self.dialect) + resolved_dialect = self._build_dialect(dialect_override) statement_expression = self._build_final_expression(copy=True) if self.enable_optimization and isinstance(statement_expression, exp.Expr): statement_expression = self._optimize_expression(statement_expression) + statement_expression = self._prepare_dialect_expression( + statement_expression, resolved_dialect, dialect_override + ) + if statement_expression.find(exp.Lock): register_lock_generator(resolved_dialect) if self._is_oracle_dialect(resolved_dialect): @@ -842,6 +937,8 @@ def _statement_from_cache_entry(self, cache_entry: "_BuilderCacheEntry", config: kwargs, parameters = self._statement_parameters(self._parameters.copy()) statement_config = config + if statement_config is not None and statement_config.dialect != cache_entry.dialect: + statement_config = statement_config.replace(dialect=cache_entry.dialect) if statement_config is None: statement_config = StatementConfig( parameter_config=ParameterStyleConfig( diff --git a/sqlspec/builder/_locking.py b/sqlspec/builder/_locking.py index 399262d31..3c4ce0d4a 100644 --- a/sqlspec/builder/_locking.py +++ b/sqlspec/builder/_locking.py @@ -38,7 +38,7 @@ def _render_lock_targets(generator: "Generator", expressions: "Iterable[exp.Expr def _lock_sql(generator: "Generator", expression: exp.Lock) -> str: - if not generator.LOCKING_READS_SUPPORTED: + if not generator.LOCKING_READS_SUPPORTED and type(generator.dialect).__name__ != "Spanner": generator.unsupported("Locking reads using 'FOR UPDATE/SHARE' are not supported") return "" @@ -46,6 +46,9 @@ def _lock_sql(generator: "Generator", expression: exp.Lock) -> str: key = expression.args.get("key") lock_type = ("FOR NO KEY UPDATE" if key else "FOR UPDATE") if update else "FOR KEY SHARE" if key else "FOR SHARE" + if expression.args.get("sqlspec_share_mode"): + lock_type = "LOCK IN SHARE MODE" + targets = _render_lock_targets(generator, expression.expressions) target_sql = f" OF {targets}" if targets else "" wait = expression.args.get("wait") diff --git a/sqlspec/builder/_select.py b/sqlspec/builder/_select.py index 8fa6b52ed..0094babbf 100644 --- a/sqlspec/builder/_select.py +++ b/sqlspec/builder/_select.py @@ -1511,7 +1511,7 @@ def for_no_key_update(self) -> "Self": assert self._expression is not None select_expr = cast("exp.Select", self._expression) - lock = exp.Lock(update=True, key=False) + lock = exp.Lock(update=True, key=True) current_locks = select_expr.args.get("locks", []) current_locks.append(lock) diff --git a/sqlspec/core/hashing.py b/sqlspec/core/hashing.py index de046b8bb..1785d61d3 100644 --- a/sqlspec/core/hashing.py +++ b/sqlspec/core/hashing.py @@ -257,7 +257,7 @@ def _expression_cache_fingerprint( settings: Any = None, ) -> str: components = ( - hash(expr), + hash_expression(expr), parameter_signature, str(dialect) if dialect is not None else "default", _freeze_cache_value(schema), diff --git a/sqlspec/data_dictionary/_types.py b/sqlspec/data_dictionary/_types.py index 4fdbfc990..d3d54ed5f 100644 --- a/sqlspec/data_dictionary/_types.py +++ b/sqlspec/data_dictionary/_types.py @@ -1422,6 +1422,7 @@ class FeatureFlags(TypedDict, total=False): supports_interleaved_tables: bool supports_json: bool supports_maps: bool + supports_on_conflict: bool supports_partitioning: bool supports_prepared_statements: bool supports_resource_groups: bool diff --git a/sqlspec/data_dictionary/dialects/bigquery/config.py b/sqlspec/data_dictionary/dialects/bigquery/config.py index fc556efe3..f315536a3 100644 --- a/sqlspec/data_dictionary/dialects/bigquery/config.py +++ b/sqlspec/data_dictionary/dialects/bigquery/config.py @@ -26,6 +26,7 @@ "supports_uuid": False, "supports_for_update": False, "supports_skip_locked": False, + "supports_on_conflict": False, } BIGQUERY_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/cockroachdb/config.py b/sqlspec/data_dictionary/dialects/cockroachdb/config.py index 254803faf..cb94ecbc5 100644 --- a/sqlspec/data_dictionary/dialects/cockroachdb/config.py +++ b/sqlspec/data_dictionary/dialects/cockroachdb/config.py @@ -24,6 +24,7 @@ "supports_for_update": True, "supports_skip_locked": True, "supports_crdb_internal_metadata": False, + "supports_on_conflict": True, } COCKROACHDB_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/duckdb/config.py b/sqlspec/data_dictionary/dialects/duckdb/config.py index 82b940e47..a180479a9 100644 --- a/sqlspec/data_dictionary/dialects/duckdb/config.py +++ b/sqlspec/data_dictionary/dialects/duckdb/config.py @@ -22,6 +22,7 @@ "supports_uuid": True, "supports_for_update": False, "supports_skip_locked": False, + "supports_on_conflict": True, } DUCKDB_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/mssql/config.py b/sqlspec/data_dictionary/dialects/mssql/config.py index 40a1dd365..35563fa89 100644 --- a/sqlspec/data_dictionary/dialects/mssql/config.py +++ b/sqlspec/data_dictionary/dialects/mssql/config.py @@ -123,6 +123,7 @@ "supports_in_memory": True, "supports_for_update": False, "supports_skip_locked": False, + "supports_on_conflict": False, } MSSQL_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/mysql/config.py b/sqlspec/data_dictionary/dialects/mysql/config.py index 482456e8d..498fdfaa7 100644 --- a/sqlspec/data_dictionary/dialects/mysql/config.py +++ b/sqlspec/data_dictionary/dialects/mysql/config.py @@ -61,6 +61,7 @@ "supports_for_update": True, "supports_sequences": False, "supports_system_versioned_tables": False, + "supports_on_conflict": False, } MYSQL_TYPE_MAPPINGS: dict[str, str] = { @@ -113,6 +114,7 @@ "supports_invisible_columns": False, "supports_invisible_indexes": False, "supports_resource_groups": False, + "supports_on_conflict": False, } MARIADB_CONFIG = DialectConfig( diff --git a/sqlspec/data_dictionary/dialects/oracle/config.py b/sqlspec/data_dictionary/dialects/oracle/config.py index b0ca053f5..a87a8006c 100644 --- a/sqlspec/data_dictionary/dialects/oracle/config.py +++ b/sqlspec/data_dictionary/dialects/oracle/config.py @@ -53,6 +53,7 @@ "supports_in_memory": True, "supports_for_update": True, "supports_skip_locked": True, + "supports_on_conflict": False, } ORACLE_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/postgres/config.py b/sqlspec/data_dictionary/dialects/postgres/config.py index 2fe1f9bf0..f3b01942b 100644 --- a/sqlspec/data_dictionary/dialects/postgres/config.py +++ b/sqlspec/data_dictionary/dialects/postgres/config.py @@ -25,6 +25,7 @@ "supports_prepared_statements": True, "supports_schemas": True, "supports_for_update": True, + "supports_on_conflict": True, } POSTGRES_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/spanner/config.py b/sqlspec/data_dictionary/dialects/spanner/config.py index 66d6b7594..8b63db7f6 100644 --- a/sqlspec/data_dictionary/dialects/spanner/config.py +++ b/sqlspec/data_dictionary/dialects/spanner/config.py @@ -13,6 +13,7 @@ "supports_interleaved_tables": True, "supports_for_update": False, "supports_skip_locked": False, + "supports_on_conflict": True, } SPANNER_TYPE_MAPPINGS: dict[str, str] = { diff --git a/sqlspec/data_dictionary/dialects/sqlite/config.py b/sqlspec/data_dictionary/dialects/sqlite/config.py index 232aaf091..466e764cd 100644 --- a/sqlspec/data_dictionary/dialects/sqlite/config.py +++ b/sqlspec/data_dictionary/dialects/sqlite/config.py @@ -23,6 +23,7 @@ "supports_uuid": False, "supports_for_update": False, "supports_skip_locked": False, + "supports_on_conflict": True, } SQLITE_TYPE_MAPPINGS: dict[str, str] = { diff --git a/tests/integration/adapters/_shared/adbc_driver.py b/tests/integration/adapters/_shared/adbc_driver.py index d09bd789c..225be0947 100644 --- a/tests/integration/adapters/_shared/adbc_driver.py +++ b/tests/integration/adapters/_shared/adbc_driver.py @@ -3,13 +3,14 @@ The contract suite owns CRUD, parameter styles, execute_many, execute_script, sequential StatementStack execution, SQLResult helpers, mapped errors, bulk operations, and multi-backend consistency. This module keeps ADBC-specific -StatementStack continue-on-error recovery and exact SQLite lock SQL generation. +StatementStack continue-on-error recovery and SQLite locking rejection. """ import pytest -from sqlspec import SQLResult, StatementStack, sql +from sqlspec import StatementStack, sql from sqlspec.adapters.adbc import AdbcDriver +from sqlspec.exceptions import SQLBuilderError from tests.conftest import requires_interpreted @@ -38,49 +39,13 @@ def test_adbc_postgresql_statement_stack_continue_on_error(adbc_postgresql_sessi @pytest.mark.xdist_group("sqlite") @pytest.mark.adbc -def test_adbc_for_update_generates_sql(adbc_sqlite_session: AdbcDriver) -> None: - """SQLite-backed ADBC strips unsupported FOR UPDATE while preserving the query.""" - adbc_sqlite_session.execute("INSERT INTO test_table_adbc (name, value) VALUES (?, ?)", ("adbc_lock", 100)) - - query = sql.select("*").from_("test_table_adbc").where_eq("name", "adbc_lock").for_update() - stmt = query.build() - - assert "FOR UPDATE" not in stmt.sql - assert "SELECT" in stmt.sql - - result = adbc_sqlite_session.execute(query) - assert isinstance(result, SQLResult) - assert result.get_data()[0]["name"] == "adbc_lock" - - -@pytest.mark.xdist_group("sqlite") -@pytest.mark.adbc -def test_adbc_for_share_generates_sql(adbc_sqlite_session: AdbcDriver) -> None: - """SQLite-backed ADBC strips unsupported FOR SHARE while preserving the query.""" - adbc_sqlite_session.execute("INSERT INTO test_table_adbc (name, value) VALUES (?, ?)", ("adbc_share", 200)) - - query = sql.select("*").from_("test_table_adbc").where_eq("name", "adbc_share").for_share() - stmt = query.build() - - assert "FOR SHARE" not in stmt.sql - assert "SELECT" in stmt.sql - - result = adbc_sqlite_session.execute(query) - assert isinstance(result, SQLResult) - assert result.get_data()[0]["name"] == "adbc_share" - - -@pytest.mark.xdist_group("sqlite") -@pytest.mark.adbc -def test_adbc_for_update_skip_locked_generates_sql(adbc_sqlite_session: AdbcDriver) -> None: - """SQLite-backed ADBC can compile a SKIP LOCKED builder path without backend locking support.""" - adbc_sqlite_session.execute("INSERT INTO test_table_adbc (name, value) VALUES (?, ?)", ("adbc_skip", 300)) - - query = sql.select("*").from_("test_table_adbc").where_eq("name", "adbc_skip").for_update(skip_locked=True) - stmt = query.build() - - assert stmt.sql is not None - - result = adbc_sqlite_session.execute(query) - assert isinstance(result, SQLResult) - assert result.get_data()[0]["name"] == "adbc_skip" +@pytest.mark.parametrize("lock_method", ["for_update", "for_share", "for_update_skip_locked"]) +def test_adbc_unsupported_lock_raises(adbc_sqlite_session: AdbcDriver, lock_method: str) -> None: + """SQLite-backed ADBC rejects locking instead of executing an unlocked query.""" + query = sql.select("*").from_("test_table_adbc") + if lock_method == "for_share": + query = query.for_share() + else: + query = query.for_update(skip_locked=lock_method == "for_update_skip_locked") + with pytest.raises(SQLBuilderError, match="does not support FOR UPDATE / row locking"): + adbc_sqlite_session.execute(query) diff --git a/tests/integration/adapters/_shared/store_behaviors.py b/tests/integration/adapters/_shared/store_behaviors.py index a49558cd3..4fb92b628 100644 --- a/tests/integration/adapters/_shared/store_behaviors.py +++ b/tests/integration/adapters/_shared/store_behaviors.py @@ -44,8 +44,10 @@ async def assert_store_delete_nonexistent_contract(store: Any) -> None: async def assert_store_expiration_int_contract(store: Any) -> None: """Stores expire entries after an integer-second TTL.""" - await store.set("expiring_session", b"data", expires_in=1) + await store.set("expiring_session", b"data", expires_in=60) assert await store.exists("expiring_session") + # The short TTL may elapse during a slow write; only require the expired state afterwards. + await store.set("expiring_session", b"data", expires_in=1) await asyncio.sleep(1.1) assert await store.get("expiring_session") is None assert not await store.exists("expiring_session") @@ -53,8 +55,9 @@ async def assert_store_expiration_int_contract(store: Any) -> None: async def assert_store_expiration_timedelta_contract(store: Any) -> None: """Stores expire entries after a timedelta TTL.""" - await store.set("expiring_session", b"data", expires_in=timedelta(seconds=1)) + await store.set("expiring_session", b"data", expires_in=timedelta(seconds=60)) assert await store.exists("expiring_session") + await store.set("expiring_session", b"data", expires_in=timedelta(seconds=1)) await asyncio.sleep(1.1) assert await store.get("expiring_session") is None diff --git a/tests/integration/adapters/duckdb/duckdb/test_driver.py b/tests/integration/adapters/duckdb/duckdb/test_driver.py index 1143b847a..e1b458397 100644 --- a/tests/integration/adapters/duckdb/duckdb/test_driver.py +++ b/tests/integration/adapters/duckdb/duckdb/test_driver.py @@ -6,6 +6,7 @@ from sqlspec import SQLResult, StatementStack, sql from sqlspec.adapters.duckdb import DuckDBDriver +from sqlspec.exceptions import SQLBuilderError from tests.conftest import requires_interpreted pytestmark = pytest.mark.xdist_group("duckdb") @@ -174,108 +175,16 @@ def test_duckdb_error_handling_and_edge_cases(duckdb_session: DuckDBDriver) -> N duckdb_session.execute_script("DROP TABLE constraint_test") -def test_duckdb_for_update_locking(duckdb_session: DuckDBDriver) -> None: - """Test FOR UPDATE row locking with DuckDB (may have limited support).""" - - # Setup test table - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") - duckdb_session.execute_script(""" - CREATE TABLE test_table ( - id INTEGER PRIMARY KEY, - name VARCHAR, - value INTEGER - ) - """) - - # Insert test data - duckdb_session.execute("INSERT INTO test_table (id, name, value) VALUES (?, ?, ?)", (1, "duckdb_lock", 100)) - - try: - duckdb_session.begin() - - # Test basic FOR UPDATE (DuckDB may have limited or no support) - result = duckdb_session.select_one( - sql.select("id", "name", "value").from_("test_table").where_eq("name", "duckdb_lock").for_update() - ) - assert result is not None - assert result["name"] == "duckdb_lock" - assert result["value"] == 100 - - duckdb_session.commit() - except Exception: - duckdb_session.rollback() - raise - finally: - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") - - -def test_duckdb_for_update_nowait(duckdb_session: DuckDBDriver) -> None: - """Test FOR UPDATE NOWAIT with DuckDB.""" - - # Setup test table - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") - duckdb_session.execute_script(""" - CREATE TABLE test_table ( - id INTEGER PRIMARY KEY, - name VARCHAR, - value INTEGER - ) - """) - - # Insert test data - duckdb_session.execute("INSERT INTO test_table (id, name, value) VALUES (?, ?, ?)", (1, "duckdb_nowait", 200)) - - try: - duckdb_session.begin() - - # Test FOR UPDATE NOWAIT - result = duckdb_session.select_one( - sql.select("*").from_("test_table").where_eq("name", "duckdb_nowait").for_update(nowait=True) - ) - assert result is not None - assert result["name"] == "duckdb_nowait" - - duckdb_session.commit() - except Exception: - duckdb_session.rollback() - raise - finally: - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") - - -def test_duckdb_for_share_locking(duckdb_session: DuckDBDriver) -> None: - """Test FOR SHARE row locking with DuckDB.""" - - # Setup test table - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") - duckdb_session.execute_script(""" - CREATE TABLE test_table ( - id INTEGER PRIMARY KEY, - name VARCHAR, - value INTEGER - ) - """) - - # Insert test data - duckdb_session.execute("INSERT INTO test_table (id, name, value) VALUES (?, ?, ?)", (1, "duckdb_share", 300)) - - try: - duckdb_session.begin() - - # Test FOR SHARE (DuckDB support may vary) - result = duckdb_session.select_one( - sql.select("id", "name", "value").from_("test_table").where_eq("name", "duckdb_share").for_share() - ) - assert result is not None - assert result["name"] == "duckdb_share" - assert result["value"] == 300 - - duckdb_session.commit() - except Exception: - duckdb_session.rollback() - raise - finally: - duckdb_session.execute_script("DROP TABLE IF EXISTS test_table") +@pytest.mark.parametrize("lock_method", ["for_update", "for_share", "for_update_nowait"]) +def test_duckdb_unsupported_lock_raises(duckdb_session: DuckDBDriver, lock_method: str) -> None: + """DuckDB rejects locking instead of executing an unlocked query.""" + query = sql.select("*").from_("test_table") + if lock_method == "for_share": + query = query.for_share() + else: + query = query.for_update(nowait=lock_method == "for_update_nowait") + with pytest.raises(SQLBuilderError, match="does not support FOR UPDATE / row locking"): + duckdb_session.select_one(query) @requires_interpreted diff --git a/tests/integration/adapters/oracle/oracledb/test_driver.py b/tests/integration/adapters/oracle/oracledb/test_driver.py index 8ad38b62e..90a4c5344 100644 --- a/tests/integration/adapters/oracle/oracledb/test_driver.py +++ b/tests/integration/adapters/oracle/oracledb/test_driver.py @@ -15,7 +15,7 @@ OracleSyncConfig, OracleSyncDriver, ) -from sqlspec.exceptions import SQLSpecError +from sqlspec.exceptions import SQLBuilderError pytestmark = pytest.mark.xdist_group("oracle") @@ -102,7 +102,7 @@ async def test_for_share_locking_unsupported(oracle_family_session: OracleFamily ) try: await _invoke(_method(oracle_family_session, "begin")) - with pytest.raises(SQLSpecError, match=r"ORA-02000.*missing COMPRESS or UPDATE keyword"): + with pytest.raises(SQLBuilderError, match="does not support FOR SHARE"): await _invoke( _method(oracle_family_session, "select_one"), sql.select("id", "name", "value").from_(table).where_eq("name", value).for_share(), diff --git a/tests/integration/adapters/sqlite/adbc/test_driver.py b/tests/integration/adapters/sqlite/adbc/test_driver.py index 65e5c1e9c..39bbc665e 100644 --- a/tests/integration/adapters/sqlite/adbc/test_driver.py +++ b/tests/integration/adapters/sqlite/adbc/test_driver.py @@ -5,17 +5,11 @@ from sqlspec.adapters.adbc import AdbcDriver from tests.integration.adapters._shared.adbc_backends import sqlite_session, test_sqlite_adbc_specific_features from tests.integration.adapters._shared.adbc_connection import test_sqlite_connection -from tests.integration.adapters._shared.adbc_driver import ( - test_adbc_for_share_generates_sql, - test_adbc_for_update_generates_sql, - test_adbc_for_update_skip_locked_generates_sql, -) +from tests.integration.adapters._shared.adbc_driver import test_adbc_unsupported_lock_raises __all__ = ( "sqlite_session", - "test_adbc_for_share_generates_sql", - "test_adbc_for_update_generates_sql", - "test_adbc_for_update_skip_locked_generates_sql", + "test_adbc_unsupported_lock_raises", "test_sqlite_adbc_specific_features", "test_sqlite_connection", ) diff --git a/tests/integration/adapters/sqlite/test_features.py b/tests/integration/adapters/sqlite/test_features.py index 113277e1a..f9fd5b85d 100644 --- a/tests/integration/adapters/sqlite/test_features.py +++ b/tests/integration/adapters/sqlite/test_features.py @@ -12,6 +12,7 @@ from sqlspec.adapters.aiosqlite import core as aiosqlite_core from sqlspec.adapters.sqlite import SqliteConfig, SqliteDriver from sqlspec.adapters.sqlite import core as sqlite_core +from sqlspec.exceptions import SQLBuilderError from tests.conftest import requires_interpreted pytestmark = pytest.mark.xdist_group("sqlite") @@ -299,24 +300,12 @@ async def test_schema_operations(sqlite_family_driver: "SQLiteFamilyDriver") -> @pytest.mark.parametrize( - ("lock_method", "lock_kwargs", "unsupported_clause"), - ( - ("for_update", {}, "FOR UPDATE"), - ("for_share", {}, "FOR SHARE"), - ("for_update", {"skip_locked": True}, "FOR UPDATE"), - ), + ("lock_method", "lock_kwargs"), (("for_update", {}), ("for_share", {}), ("for_update", {"skip_locked": True})) ) -async def test_unsupported_lock_clause_is_stripped( - sqlite_family_driver: "SQLiteFamilyDriver", lock_method: str, lock_kwargs: "dict[str, Any]", unsupported_clause: str +async def test_unsupported_lock_clause_raises( + sqlite_family_driver: "SQLiteFamilyDriver", lock_method: str, lock_kwargs: "dict[str, Any]" ) -> None: - prefix = _prefix(sqlite_family_driver) - name = f"{prefix}-{lock_method}" - await _invoke( - _method(sqlite_family_driver, "execute"), "INSERT INTO test_table (name, value) VALUES (?, ?)", (name, 100) - ) - query = sql.select("*").from_("test_table").where_eq("name", name) + query = sql.select("*").from_("test_table") query = cast("Any", getattr(query, lock_method))(**lock_kwargs) - statement = query.build() - assert unsupported_clause not in statement.sql - result = await _invoke(_method(sqlite_family_driver, "execute"), query) - assert result.get_data()[0]["name"] == name + with pytest.raises(SQLBuilderError, match="does not support FOR UPDATE / row locking"): + await _invoke(_method(sqlite_family_driver, "execute"), query) diff --git a/tests/unit/builder/test_dialect_override.py b/tests/unit/builder/test_dialect_override.py index 5675b2b4a..0c512211c 100644 --- a/tests/unit/builder/test_dialect_override.py +++ b/tests/unit/builder/test_dialect_override.py @@ -183,3 +183,35 @@ def test_to_sql_dialect_override_with_complex_query() -> None: assert "JOIN" in mysql_sql assert "WHERE" in postgres_sql assert "WHERE" in mysql_sql + + +def test_mssql_alias_builds_tsql() -> None: + """Test dialect override 'mssql' normalizes to 'tsql' and renders identical T-SQL.""" + query = sql.select("id", "name").from_("products").limit(10) + mssql_sql = query.build(dialect="mssql").sql + tsql_sql = query.build(dialect="tsql").sql + assert "TOP" in mssql_sql + assert mssql_sql == tsql_sql + + +def test_mariadb_and_cockroachdb_dialect_aliases() -> None: + """Test dialect aliases for mariadb and cockroachdb render expected syntax.""" + query = sql.select("id", "name").from_("products") + mariadb_sql = query.build(dialect="mariadb").sql + mysql_sql = query.build(dialect="mysql").sql + assert mariadb_sql == mysql_sql + + cockroach_query = sql.select("id", "name").from_("products").for_update(skip_locked=True) + cockroach_sql = cockroach_query.build(dialect="cockroachdb").sql + postgres_sql = cockroach_query.build(dialect="postgres").sql + assert "FOR UPDATE SKIP LOCKED" in cockroach_sql + assert cockroach_sql == postgres_sql + + +def test_dialect_class_override_and_alias_metadata() -> None: + from sqlglot.dialects.mysql import MySQL + + query = sql.select("id").from_("products").limit(10) + assert query.build(dialect=MySQL).sql == query.build(dialect="mysql").sql + assert query.build(dialect=MySQL()).sql == query.build(dialect="mysql").sql + assert query.build(dialect="mssql").dialect == "tsql" diff --git a/tests/unit/builder/test_insert_builder.py b/tests/unit/builder/test_insert_builder.py index 6ab2c0509..dcdb6e27d 100644 --- a/tests/unit/builder/test_insert_builder.py +++ b/tests/unit/builder/test_insert_builder.py @@ -383,3 +383,95 @@ def test_on_conflict_with_values_from() -> None: assert "DO UPDATE" in stmt.sql assert "name_1" in stmt.parameters assert stmt.parameters["name_1"] == "Updated" + + +def test_on_conflict_mysql_transpiles() -> None: + """Test ON CONFLICT transpiles to ON DUPLICATE KEY UPDATE for MySQL and MariaDB.""" + query = sql.insert("users").values(id=1, name="John").on_conflict("id").do_update(name="Updated") + mysql_stmt = query.build(dialect="mysql") + assert "ON DUPLICATE KEY UPDATE" in mysql_stmt.sql + assert "ON CONFLICT" not in mysql_stmt.sql + assert "name" in mysql_stmt.sql + + mariadb_stmt = query.build(dialect="mariadb") + assert "ON DUPLICATE KEY UPDATE" in mariadb_stmt.sql + assert "ON CONFLICT" not in mariadb_stmt.sql + + nothing_query = sql.insert("users").values(id=1, name="John").on_conflict("id").do_nothing() + nothing_mysql = nothing_query.build(dialect="mysql") + assert "ON DUPLICATE KEY UPDATE" in nothing_mysql.sql + assert "id = id" in nothing_mysql.sql or "`id` = `id`" in nothing_mysql.sql + + +@pytest.mark.parametrize("dialect", ["oracle", "tsql", "mssql", "bigquery"]) +def test_on_conflict_raises_oracle_tsql(dialect: str) -> None: + """Test ON CONFLICT raises SQLBuilderError mentioning sql.merge() on unsupported dialects.""" + query = sql.insert("users").values(id=1, name="John").on_conflict("id").do_update(name="Updated") + with pytest.raises(SQLBuilderError, match=r"sql\.merge\(\)"): + query.build(dialect=dialect) + + +def test_on_conflict_postgres_unchanged() -> None: + """Test ON CONFLICT on postgres retains standard ON CONFLICT clause.""" + query = sql.insert("users").values(id=1, name="John").on_conflict("id").do_update(name="Updated") + stmt = query.build(dialect="postgres") + assert "ON CONFLICT" in stmt.sql + assert "DO UPDATE" in stmt.sql + assert "ON DUPLICATE KEY UPDATE" not in stmt.sql + + +def test_to_statement_translates_conflict_without_mutating_builder() -> None: + from sqlspec.core import StatementConfig + + query = ( + sql + .insert("users") + .values(id=1, name="John") + .on_conflict("id") + .do_update(name=exp.column("name", table="excluded")) + ) + original = query.build(dialect="postgres").sql + statement = query.to_statement(StatementConfig(dialect="mysql")) + assert "ON DUPLICATE KEY UPDATE" in statement.sql + assert "VALUES(" in statement.sql + assert "excluded" not in statement.sql.lower() + assert query.build(dialect="postgres").sql == original + + +def test_to_statement_rejects_unsupported_conflict() -> None: + from sqlspec.core import StatementConfig + + query = sql.insert("users").values(id=1).on_conflict("id").do_nothing() + with pytest.raises(SQLBuilderError, match=r"sql\.merge\(\)"): + query.to_statement(StatementConfig(dialect="oracle")) + + +def test_mysql_do_nothing_requires_known_column() -> None: + query = sql.insert("users").values(1).on_conflict().do_nothing() + with pytest.raises(SQLBuilderError, match="requires a conflict column"): + query.build(dialect="mysql") + + +@pytest.mark.parametrize("argument", ["where", "index_predicate", "constraint"]) +def test_mysql_rejects_conflict_semantics_it_cannot_preserve(argument: str) -> None: + query = sql.insert("users").values(id=1).on_conflict("id").do_update(id=2) + conflict = query.get_insert_expression().args["conflict"] + conflict.set(argument, exp.to_identifier("restricted")) + with pytest.raises(SQLBuilderError, match="cannot preserve"): + query.build(dialect="mysql") + + +@pytest.mark.parametrize("dialect", ["spanner", "spangres"]) +def test_spanner_native_conflicts(dialect: str) -> None: + from sqlspec.core import StatementConfig + + query = sql.insert("users").values(id=1).on_conflict("id").do_nothing() + assert "ON CONFLICT" in query.to_statement(StatementConfig(dialect=dialect)).sql + query = sql.insert("users").values(id=1).on_conflict("id").do_update(id=exp.column("id", table="excluded")) + assert "DO UPDATE" in query.build(dialect=dialect).sql + + +def test_spangres_rejects_non_insert_value_conflict_update() -> None: + query = sql.insert("users").values(id=1).on_conflict("id").do_update(id=2) + with pytest.raises(SQLBuilderError, match="require excluded column values"): + query.build(dialect="spangres") diff --git a/tests/unit/builder/test_select_locking.py b/tests/unit/builder/test_select_locking.py index 0f1ecc6be..f0805e716 100644 --- a/tests/unit/builder/test_select_locking.py +++ b/tests/unit/builder/test_select_locking.py @@ -1,8 +1,11 @@ """Unit tests for SELECT locking functionality (FOR UPDATE, FOR SHARE, etc).""" +import re + import pytest from sqlspec import sql +from sqlspec.data_dictionary import DialectConfig, register_dialect from sqlspec.exceptions import SQLBuilderError @@ -234,5 +237,96 @@ def test_complex_join_with_for_update_of() -> None: # Both tables should be mentioned in the OF clause assert "j" in sql_content assert "u" in sql_content - # Should contain the companies reference as well assert "c" in sql_content + + +@pytest.mark.parametrize("dialect", ["tsql", "mssql", "sqlite", "duckdb", "bigquery"]) +@pytest.mark.parametrize("statement", [False, True]) +def test_for_update_raises_on_unsupported_dialects(dialect: str, statement: bool) -> None: + """Test FOR UPDATE raises SQLBuilderError on dialects without row lock support.""" + from sqlspec.core import StatementConfig + + query = sql.select("*").from_("job").for_update() + with pytest.raises(SQLBuilderError, match="does not support FOR UPDATE"): + if statement: + query.to_statement(StatementConfig(dialect=dialect)) + else: + query.build(dialect=dialect) + + +@pytest.mark.parametrize("dialect", ["postgres", "mysql", "oracle", "cockroachdb", "mariadb", "spanner", "spangres"]) +def test_for_update_supported_dialects(dialect: str) -> None: + """Test FOR UPDATE renders on dialects supporting row locks.""" + query = sql.select("*").from_("job").for_update() + stmt = query.build(dialect=dialect) + assert "FOR UPDATE" in stmt.sql + + +def test_skip_locked_flag_gate() -> None: + """Test SKIP LOCKED raises SQLBuilderError when supports_skip_locked is False.""" + custom_config = DialectConfig( + name="custom_no_skip_locked", + feature_flags={"supports_for_update": True, "supports_skip_locked": False}, + feature_versions={}, + type_mappings={}, + version_pattern=re.compile(r".*"), + ) + register_dialect(custom_config) + + query = sql.select("*").from_("job").for_update(skip_locked=True) + with pytest.raises(SQLBuilderError, match="does not support SKIP LOCKED"): + query.build(dialect="custom_no_skip_locked") + + +@pytest.mark.parametrize("statement", [False, True]) +def test_oracle_for_share_rejected(statement: bool) -> None: + from sqlspec.core import StatementConfig + + query = sql.select("id").from_("job").for_share() + with pytest.raises(SQLBuilderError, match="does not support FOR SHARE"): + if statement: + query.to_statement(StatementConfig(dialect="oracle")) + else: + query.build(dialect="oracle") + + +@pytest.mark.parametrize("statement", [False, True]) +def test_mariadb_shared_lock_rendering(statement: bool) -> None: + from sqlspec.core import StatementConfig + + query = sql.select("id").from_("job").for_share(skip_locked=True) + rendered = query.to_statement(StatementConfig(dialect="mariadb")) if statement else query.build(dialect="mariadb") + assert "LOCK IN SHARE MODE SKIP LOCKED" in rendered.sql + assert "FOR SHARE" in query.build(dialect="postgres").sql + + +@pytest.mark.parametrize("dialect", ["spanner", "spangres"]) +def test_spanner_statement_locking(dialect: str) -> None: + from sqlspec.core import StatementConfig + + query = sql.select("id").from_("job").for_update() + assert "FOR UPDATE" in query.to_statement(StatementConfig(dialect=dialect)).sql + for locked_query in ( + sql.select("id").from_("job").for_update(skip_locked=True), + sql.select("id").from_("job").for_update(nowait=True), + sql.select("id").from_("job").for_update(of="job"), + sql.select("id").from_("job").for_share(), + ): + with pytest.raises(SQLBuilderError, match="only plain FOR UPDATE"): + locked_query.build(dialect=dialect) + + +@pytest.mark.parametrize("dialect", ["mysql", "mariadb", "oracle"]) +def test_postgresql_key_lock_modes_rejected_elsewhere(dialect: str) -> None: + for query in (sql.select("id").from_("job").for_no_key_update(), sql.select("id").from_("job").for_key_share()): + with pytest.raises(SQLBuilderError, match="PostgreSQL key lock modes"): + query.build(dialect=dialect) + + +def test_mariadb_rejects_lock_targets() -> None: + for query in ( + sql.select("id").from_("job").for_update(of="job"), + sql.select("id").from_("job").for_share(of="job"), + ): + with pytest.raises(SQLBuilderError, match="do not support OF targets"): + query.build(dialect="mariadb") diff --git a/tests/unit/builder/test_sqlglot_arg_contracts.py b/tests/unit/builder/test_sqlglot_arg_contracts.py index 466e1c3f1..e8c395d55 100644 --- a/tests/unit/builder/test_sqlglot_arg_contracts.py +++ b/tests/unit/builder/test_sqlglot_arg_contracts.py @@ -31,6 +31,7 @@ ("sqlspec/adapters/duckdb/core.py", "part", "quoted"), ("sqlspec/builder/_base.py", "final_expression", "with_"), ("sqlspec/builder/_base.py", "expression", "conflict"), + ("sqlspec/builder/_base.py", "lock", "sqlspec_share_mode"), ("sqlspec/builder/_base.py", "node", "quoted"), ("sqlspec/builder/_base.py", "optimized", "conflict"), ("sqlspec/builder/_dml.py", "current_expr", "expression"),