From 10ba2fceb969f7c7f085f174a6e101ec266c3131 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 20:19:14 +0000 Subject: [PATCH 01/15] feat(spanner): transaction retry, commit inlining, multiplexed pooling, FLOAT32 vectors, and partitioned DML - Implement run_in_transaction retry loop with exponential backoff on Aborted (409) - Support last_statement=True in transaction commits to eliminate a network round trip - Support multiplexed session pooling and PingingPool keepalive intervals - Support FLOAT32 vector search parameter coercion, Decimal type inference, and INTERVAL types - Optimize JsonObject deserialization without double string conversion - Add execute_partitioned_dml and use for bulk table truncation in load_from_arrow(overwrite=True) - Decouple enable_batch_write_api from requiring an active transaction - Forward QueryOptions, RequestOptions, and DirectedReadOptions per session and statement --- sqlspec/adapters/spanner/adk/store.py | 91 ++------- sqlspec/adapters/spanner/config.py | 182 ++++++++++++++--- sqlspec/adapters/spanner/core.py | 43 +++- sqlspec/adapters/spanner/driver.py | 193 +++++++++++++++--- sqlspec/adapters/spanner/litestar/store.py | 62 ++---- sqlspec/adapters/spanner/type_converter.py | 113 ++++++++-- sqlspec/core/parameters/_types.py | 7 +- sqlspec/protocols.py | 3 + .../test_spanner/test_batch_write_api.py | 9 +- .../test_spanner/test_litestar_store.py | 28 +-- .../test_load_from_arrow_mutations.py | 9 +- .../test_spanner_arrow_overwrite.py | 35 ++++ .../test_spanner/test_spanner_batch_write.py | 59 ++++++ .../test_spanner/test_spanner_json.py | 54 +++++ .../test_spanner_last_statement.py | 58 ++++++ .../test_spanner_partitioned_dml.py | 146 +++++++++++++ .../test_spanner/test_spanner_pinging_pool.py | 51 +++++ .../test_spanner/test_spanner_pool.py | 71 +++++++ .../test_spanner_query_options.py | 101 +++++++++ .../test_spanner_request_options.py | 105 ++++++++++ .../test_spanner/test_spanner_stores.py | 66 ++++++ .../test_spanner/test_spanner_transaction.py | 115 +++++++++++ .../test_spanner_type_inference.py | 42 ++++ .../test_spanner/test_spanner_vector.py | 48 +++++ 24 files changed, 1462 insertions(+), 229 deletions(-) create mode 100644 tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_batch_write.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_json.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_last_statement.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_pool.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_query_options.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_request_options.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_stores.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_transaction.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_type_inference.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_vector.py diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index 5facfb347..9c347d8fa 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -10,7 +10,6 @@ from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.adapters.spanner.config import SpannerSyncConfig from sqlspec.config import ADKConfig -from sqlspec.exceptions import OperationalError from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore from sqlspec.protocols import SpannerParamTypesProtocol @@ -20,7 +19,7 @@ from collections.abc import Sequence from sqlspec.adapters.spanner._typing import SpannerDatabase as Database - from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction + from sqlspec.adapters.spanner.driver import SpannerSyncDriver from sqlspec.extensions.adk import SessionOrderBy, StoredMemory __all__ = ("SpannerADKConfig", "SpannerADKRetentionConfig", "SpannerSyncADKMemoryStore", "SpannerSyncADKStore") @@ -225,7 +224,11 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - self._database().run_in_transaction(_SpannerWriteJob(statements)) # type: ignore[no-untyped-call] + def _job(driver: "SpannerSyncDriver") -> None: + for sql, params, _ in statements: + driver.execute(sql, params) + + self._config.run_in_transaction(_job) def _session_param_types(self, include_owner: bool) -> "dict[str, Any]": json_type = _json_param_type() @@ -616,32 +619,29 @@ def _append_event(self, event_record: StoredEvent) -> None: def _delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._events_table} WHERE timestamp < @before" params: dict[str, Any] = {"before": before} - types: dict[str, Any] = {"before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - types["app_name"] = SPANNER_PARAM_TYPES.STRING - return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) + result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) + return int(getattr(result, "rowcount", 0)) def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._session_table} WHERE update_time < @updated_before" params: dict[str, Any] = {"updated_before": updated_before} - types: dict[str, Any] = {"updated_before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - types["app_name"] = SPANNER_PARAM_TYPES.STRING - return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) + result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) + return int(getattr(result, "rowcount", 0)) def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._user_state_table} WHERE update_time < @updated_before" params: dict[str, Any] = {"updated_before": updated_before} - types: dict[str, Any] = {"updated_before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - types["app_name"] = SPANNER_PARAM_TYPES.STRING - return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) + result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) + return int(getattr(result, "rowcount", 0)) def _get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = @app_name LIMIT 1" @@ -879,10 +879,15 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - self._database().run_in_transaction(_SpannerMemoryWriteJob(statements)) # type: ignore[no-untyped-call] + def _job(driver: "SpannerSyncDriver") -> None: + for sql, params, _ in statements: + driver.execute(sql, params) + + self._config.run_in_transaction(_job) def _execute_update(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> int: - return int(self._database().run_in_transaction(_SpannerMemoryUpdateJob(sql, params, types))) # type: ignore[no-untyped-call] + result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) + return int(getattr(result, "rowcount", 0)) def _memory_param_types(self, include_owner: bool) -> "dict[str, Any]": types: dict[str, Any] = { @@ -1209,64 +1214,6 @@ def _spanner_drop_statement_table(statement: str, existing_tables: "set[str]") - return None -class _SpannerWriteJob: - __slots__ = ("_statements",) - - def __init__(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - self._statements = statements - - def __call__(self, transaction: "Transaction") -> None: - if len(self._statements) > 1: - status, _row_counts = transaction.batch_update(self._statements) # type: ignore[no-untyped-call] - if status.code != 0: - msg = f"Spanner batch update failed (code {status.code}): {status.message}" - raise OperationalError(msg) - return - for sql, params, types in self._statements: - transaction.execute_update(sql, params=params, param_types=types) # type: ignore[no-untyped-call] - - -class _SpannerMemoryWriteJob: - __slots__ = ("_statements",) - - def __init__(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - self._statements = statements - - def __call__(self, transaction: "Transaction") -> None: - if len(self._statements) > 1: - status, _row_counts = transaction.batch_update(self._statements) # type: ignore[no-untyped-call] - if status.code != 0: - msg = f"Spanner batch update failed (code {status.code}): {status.message}" - raise OperationalError(msg) - return - for sql, params, types in self._statements: - transaction.execute_update(sql, params=params, param_types=types) # type: ignore[no-untyped-call] - - -class _SpannerUpdateJob: - __slots__ = ("_params", "_sql", "_types") - - def __init__(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: - self._sql = sql - self._params = params - self._types = types - - def __call__(self, transaction: "Transaction") -> int: - return int(transaction.execute_update(self._sql, params=self._params, param_types=self._types)) # type: ignore[no-untyped-call] - - -class _SpannerMemoryUpdateJob: - __slots__ = ("_params", "_sql", "_types") - - def __init__(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: - self._sql = sql - self._params = params - self._types = types - - def __call__(self, transaction: "Transaction") -> int: - return int(transaction.execute_update(self._sql, params=self._params, param_types=self._types)) # type: ignore[no-untyped-call] - - class _SpannerReadProtocol(Protocol): def execute_sql( self, sql: str, params: "dict[str, Any] | None" = None, param_types: "dict[str, Any] | None" = None diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index aa475eafd..8ca11ed35 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -9,8 +9,9 @@ from sqlspec.adapters.spanner._typing import SpannerTransactionType as TransactionType from sqlspec.adapters.spanner.core import apply_driver_features, default_statement_config from sqlspec.adapters.spanner.driver import SpannerSessionContext, SpannerSyncDriver +from sqlspec.adapters.spanner.type_converter import coerce_params_for_spanner, infer_spanner_param_types from sqlspec.config import SyncDatabaseConfig -from sqlspec.core import TypeCoercionCapabilities +from sqlspec.core import SQL, TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints @@ -132,7 +133,7 @@ class SpannerConnectionParams(TypedDict): class SpannerPoolParams(SpannerConnectionParams): """Session pool configuration.""" - pool_type: "NotRequired[type[AbstractSessionPool]]" + pool_type: "NotRequired[type[AbstractSessionPool] | str | None]" size: "NotRequired[int]" target_size: "NotRequired[int]" max_sessions: "NotRequired[int]" @@ -140,7 +141,9 @@ class SpannerPoolParams(SpannerConnectionParams): session_labels: "NotRequired[dict[str, str]]" labels: "NotRequired[dict[str, str]]" ping_interval: "NotRequired[int]" + ping_timeout: "NotRequired[float]" max_age_minutes: "NotRequired[int]" + enable_multiplexed_sessions: "NotRequired[bool]" class SpannerDriverFeatures(TypedDict): @@ -174,6 +177,7 @@ class SpannerDriverFeatures(TypedDict): retry: "NotRequired[Retry | None]" timeout: "NotRequired[float | None]" request_options: "NotRequired[RequestOptions | dict[str, Any] | None]" + query_options: "NotRequired[ExecuteSqlRequest.QueryOptions | dict[str, Any] | None]" directed_read_options: "NotRequired[DirectedReadOptions | None]" session_labels: "NotRequired[dict[str, str]]" enable_events: "NotRequired[bool]" @@ -212,22 +216,26 @@ def __init__(self, config: "SpannerSyncConfig", transaction: bool = False) -> No def __enter__(self) -> SpannerConnection: database = self._config.get_database() if self._transaction: - manager = cast("Any", database).sessions_manager - self._session = manager.get_session(TransactionType.READ_WRITE) - try: - txn = self._session.transaction() - txn.__enter__() - self._connection = cast("SpannerConnection", txn) - except Exception: - manager.put_session(self._session) - self._session = None - raise - else: - return self._connection - else: - self._session = cast("Any", database).snapshot(multi_use=True) - self._connection = cast("SpannerConnection", self._session.__enter__()) + manager = getattr(database, "sessions_manager", None) + if manager is not None and hasattr(manager, "get_session"): + self._session = manager.get_session(TransactionType.READ_WRITE) + try: + txn = self._session.transaction() + txn.__enter__() + self._connection = cast("SpannerConnection", txn) + except Exception: + manager.put_session(self._session) + self._session = None + raise + else: + return self._connection + txn = database.transaction() + self._session = txn + self._connection = cast("SpannerConnection", txn.__enter__()) return self._connection + self._session = cast("Any", database).snapshot(multi_use=True) + self._connection = cast("SpannerConnection", self._session.__enter__()) + return self._connection def __exit__( self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" @@ -256,7 +264,12 @@ def __exit__( txn.rollback() finally: if self._session: - cast("Any", self._config.get_database()).sessions_manager.put_session(self._session) + db = self._config.get_database() + manager = getattr(db, "sessions_manager", None) + if manager is not None and hasattr(manager, "put_session"): + manager.put_session(self._session) + elif hasattr(self._session, "__exit__"): + self._session.__exit__(exc_type, exc_val, exc_tb) elif self._session: self._session.__exit__(exc_type, exc_val, exc_tb) @@ -324,10 +337,15 @@ def __init__( ): self.connection_config["session_labels"] = legacy_session_labels - from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool - self.connection_config.setdefault("size", self.connection_config.pop("max_sessions", 10)) - self.connection_config.setdefault("pool_type", FixedSizePool) + enable_multiplexed = self.connection_config.get("enable_multiplexed_sessions", True) + if enable_multiplexed and "pool_type" not in self.connection_config: + self.connection_config["pool_type"] = None + elif not enable_multiplexed and "pool_type" not in self.connection_config: + from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool + + self.connection_config["pool_type"] = PingingPool + self.connection_config.setdefault("ping_interval", 1800) statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, raw_driver_features) @@ -362,7 +380,11 @@ def get_database(self) -> "Database": msg = "instance_id and database_id are required." raise ImproperConfigurationError(msg) - if self.connection_instance is None: + pool_type = self.connection_config.get("pool_type") + enable_multiplexed = self.connection_config.get("enable_multiplexed_sessions", True) + is_multiplexed = enable_multiplexed and (pool_type is None or pool_type == "multiplexed") + + if not is_multiplexed and self.connection_instance is None: self.connection_instance = self.provide_pool() if self._database is None: @@ -372,7 +394,8 @@ def get_database(self) -> "Database": if instance_labels is not None: instance_kwargs["labels"] = instance_labels database_kwargs = self._connection_kwargs_for(_DATABASE_CONFIG_FIELDS) - database_kwargs["pool"] = self.connection_instance + if self.connection_instance is not None: + database_kwargs["pool"] = self.connection_instance self._database = client.instance(instance_id, **instance_kwargs).database( # type: ignore[no-untyped-call] database_id, **database_kwargs ) @@ -401,11 +424,16 @@ def _create_pool(self) -> "AbstractSessionPool": msg = "instance_id and database_id are required." raise ImproperConfigurationError(msg) - pool_type = cast("type[AbstractSessionPool]", self.connection_config.get("pool_type", FixedSizePool)) + raw_pool_type = self.connection_config.get("pool_type") + if raw_pool_type is None or raw_pool_type == "multiplexed": + pool_type = PingingPool + else: + pool_type = cast("type[AbstractSessionPool]", raw_pool_type) labels = self.connection_config.get("session_labels", self.connection_config.get("labels")) pool_kwargs: dict[str, Any] = self._pool_base_kwargs(labels=cast("dict[str, str] | None", labels)) if issubclass(pool_type, PingingPool): + self.connection_config.setdefault("ping_interval", 1800) pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "ping_interval"})) elif issubclass(pool_type, FixedSizePool): pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "max_age_minutes"})) @@ -480,6 +508,7 @@ def provide_session( transaction: "bool" = _DEFAULT_SESSION_TRANSACTION, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -497,6 +526,7 @@ def provide_session( Snapshot (False). request_options: Session-scoped RequestOptions for Spanner statements. directed_read_options: Session-scoped DirectedReadOptions for reads. + query_options: Session-scoped QueryOptions for Spanner statements. retry: Session-scoped retry policy for Spanner statement calls. timeout: Session-scoped timeout for Spanner statement calls. **kwargs: Additional keyword arguments. @@ -514,6 +544,7 @@ def provide_session( driver_features=self._session_driver_features( request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, ), @@ -526,6 +557,7 @@ def provide_write_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -537,6 +569,7 @@ def provide_write_session( transaction=True, request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, **kwargs, @@ -548,6 +581,7 @@ def provide_read_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -563,26 +597,124 @@ def provide_read_session( transaction=False, request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, **kwargs, ) + def run_in_transaction(self, func: "Callable[[SpannerSyncDriver], Any]", *args: Any, **kwargs: Any) -> Any: + """Execute a unit of work inside a transaction, retrying on abort. + + Args: + func: Callback taking a prepared SpannerSyncDriver and returning a result. + *args: Positional arguments passed to the callback. + **kwargs: Keyword arguments passed to the callback. + + Returns: + The return value of the callback. + """ + database = self.get_database() + + def _callback(spanner_transaction: Any, *cb_args: Any, **cb_kwargs: Any) -> Any: + driver = SpannerSyncDriver( + connection=spanner_transaction, + statement_config=self.statement_config, + driver_features=self.driver_features, + ) + prepared_driver = self._prepare_driver(driver) + call_args = cb_args or args + call_kwargs = cb_kwargs or kwargs + try: + return func(prepared_driver, *call_args, **call_kwargs) + except Exception as exc: + cause = getattr(exc, "__cause__", None) + from google.api_core import exceptions as api_exceptions + + if isinstance(exc, api_exceptions.Aborted): + raise + if cause is not None and isinstance(cause, api_exceptions.Aborted): + raise cause from exc + raise + + return database.run_in_transaction(_callback, *args, **kwargs) + + def execute_partitioned_dml( + self, + statement: "SQL | str", + *parameters: Any, + query_options: Any = None, + request_options: Any = None, + exclude_txn_from_change_streams: bool = False, + **kwargs: Any, + ) -> int: + """Execute a Partitioned DML statement across database partitions. + + Args: + statement: The SQL string or SQL object to execute. + *parameters: Positional parameters or parameter mapping. + query_options: Optional Spanner QueryOptions. + request_options: Optional Spanner RequestOptions. + exclude_txn_from_change_streams: Whether to exclude the transaction from change streams. + **kwargs: Additional keyword arguments or parameters. + + Returns: + The number of affected rows. + """ + database = self.get_database() + if isinstance(statement, SQL): + sql_statement = statement + else: + sql_statement = SQL(statement, *parameters, statement_config=self.statement_config, **kwargs) + + sql, raw_params = sql_statement.compile() + params = raw_params if isinstance(raw_params, dict) else None + coerced_params = coerce_params_for_spanner( + params, + json_serializer=self.driver_features.get("json_serializer"), + enable_uuid_conversion=self.driver_features.get("enable_uuid_conversion", True), + ) + param_types = infer_spanner_param_types(params) + + effective_request_options = request_options or self.driver_features.get("request_options") + effective_query_options = query_options or self.driver_features.get("query_options") + + call_kwargs: dict[str, Any] = { + "params": coerced_params, + "param_types": param_types, + "exclude_txn_from_change_streams": exclude_txn_from_change_streams, + } + if effective_query_options is not None: + call_kwargs["query_options"] = effective_query_options + if effective_request_options is not None: + call_kwargs["request_options"] = effective_request_options + + return database.execute_partitioned_dml(sql, **call_kwargs) + def _session_driver_features( self, *, request_options: "RequestOptions | dict[str, Any] | None", directed_read_options: "DirectedReadOptions | None", + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None", timeout: "float | None", ) -> "dict[str, Any]": - if request_options is None and directed_read_options is None and retry is None and timeout is None: + if ( + request_options is None + and directed_read_options is None + and query_options is None + and retry is None + and timeout is None + ): return self.driver_features driver_features = dict(self.driver_features) if request_options is not None: driver_features["request_options"] = request_options if directed_read_options is not None: driver_features["directed_read_options"] = directed_read_options + if query_options is not None: + driver_features["query_options"] = query_options if retry is not None: driver_features["retry"] = retry if timeout is not None: diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index d56dbf572..eecada431 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -310,19 +310,44 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec return _create_spanner_error(error, SQLSpecError, "error") +def _unwrap_spanner_json_object(val: Any) -> Any: + """Recursively unwrap Spanner JsonObject instances into native Python primitives.""" + if isinstance(val, JsonObject): + if getattr(val, "_is_null", False): + return None + if getattr(val, "_is_array", False): + array_val = getattr(val, "_array_value", None) + return [_unwrap_spanner_json_object(item) for item in array_val] if array_val is not None else [] + if getattr(val, "_is_scalar_value", False): + return getattr(val, "_simple_value", None) + return {k: _unwrap_spanner_json_object(v) for k, v in val.items()} + if isinstance(val, dict): + return {k: _unwrap_spanner_json_object(v) for k, v in val.items()} + if isinstance(val, (list, tuple)): + return [_unwrap_spanner_json_object(item) for item in val] + return val + + def _convert_json_row_value(value: Any, *, json_deserializer: "Callable[[str], Any]") -> Any: """Convert a native Spanner JSON cell using the configured deserializer.""" if isinstance(value, JsonObject): - json_value = cast("Any", value).serialize() + if json_deserializer is from_json: + return _unwrap_spanner_json_object(value) + if getattr(value, "_is_null", False): + return None + try: + serialized = cast("Any", value).serialize() + if serialized is None: + return None + return json_deserializer(serialized) + except (TypeError, ValueError): + return _unwrap_spanner_json_object(value) elif isinstance(value, str): - json_value = value - else: - return value - - try: - return json_deserializer(json_value) - except (TypeError, ValueError): - return value + try: + return json_deserializer(value) + except (TypeError, ValueError): + return value + return value def _create_spanner_error(error: Any, error_class: type[SQLSpecError], description: str) -> SQLSpecError: diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 322426006..f78903979 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -28,6 +28,7 @@ ) from sqlspec.adapters.spanner.data_dictionary import SpannerDataDictionary from sqlspec.core import StatementConfig, register_driver_profile +from sqlspec.core.statement import SQL from sqlspec.driver import ( BaseSyncExceptionHandler, ExecutionResult, @@ -45,11 +46,11 @@ from sqlspec.adapters.spanner._typing import SpannerConnection from sqlspec.adapters.spanner._typing import SpannerDirectedReadOptions as DirectedReadOptions + from sqlspec.adapters.spanner._typing import SpannerExecuteSqlRequest as ExecuteSqlRequest from sqlspec.adapters.spanner._typing import SpannerRequestOptions as RequestOptions from sqlspec.adapters.spanner._typing import SpannerRetry as Retry from sqlspec.builder import QueryBuilder from sqlspec.core import ArrowResult, SQLResult, Statement, StatementFilter - from sqlspec.core.statement import SQL from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import SchemaT, StatementParameters @@ -93,7 +94,7 @@ class SpannerSyncDriver(SyncDriverAdapterBase): """Synchronous Spanner driver operating on Snapshot or Transaction contexts.""" dialect: "DialectType" = "spanner" - __slots__ = ("_data_dictionary", "_pending_execute_options", "_row_plan_cache", "_row_plan_deserializer") + __slots__ = ("_config", "_data_dictionary", "_pending_execute_options", "_row_plan_cache", "_row_plan_deserializer") def __init__( self, @@ -106,6 +107,7 @@ def __init__( statement_config = default_statement_config super().__init__(connection=connection, statement_config=statement_config, driver_features=features) + self._config: Any = None self._data_dictionary: SpannerDataDictionary | None = None self._pending_execute_options: _PerCallExecuteOptions | None = None self._row_plan_cache: dict[int, tuple[Any, list[str], tuple[tuple[int, Any], ...] | None]] = {} @@ -176,7 +178,7 @@ def dispatch_execute_many(self, cursor: "SpannerConnection", statement: "SQL") - _coerce = self._coerce_params _infer = self._infer_param_types - execute_kwargs = self._execute_kwargs() + execute_kwargs = self._execute_kwargs(for_batch=True) param_types_cache: dict[tuple[tuple[str, type[Any], Any], ...], dict[str, Any]] = {} empty_param_types: dict[str, Any] = {} batch_args: list[tuple[str, dict[str, Any] | None, dict[str, Any]]] = [] @@ -237,16 +239,18 @@ def begin(self) -> None: return None def commit(self) -> None: - if isinstance(self.connection, SpannerTransaction): + if isinstance(self.connection, SpannerTransaction) or supports_write(self.connection): writer = cast("_SpannerWriteProtocol", self.connection) - if writer.committed is not None: + if getattr(writer, "committed", None) is not None: return - writer.commit() + if callable(getattr(writer, "commit", None)): + writer.commit() def rollback(self) -> None: - if isinstance(self.connection, SpannerTransaction): + if isinstance(self.connection, SpannerTransaction) or supports_write(self.connection): writer = cast("_SpannerWriteProtocol", self.connection) - writer.rollback() + if callable(getattr(writer, "rollback", None)): + writer.rollback() def create_savepoint(self, name: str) -> None: """Raise because Spanner does not support savepoints. @@ -281,6 +285,119 @@ def with_cursor(self, connection: "SpannerConnection") -> "SpannerSyncCursor": def handle_database_exceptions(self) -> "SpannerExceptionHandler": return SpannerExceptionHandler() + def _get_database(self) -> Any: + if self._config is not None: + return self._config.get_database() + session = getattr(self.connection, "_session", None) + if session is not None: + database = getattr(session, "_database", None) + if database is not None: + return database + database = getattr(self.connection, "_database", None) + if database is not None: + return database + return None + + def run_in_transaction(self, func: "Callable[[SpannerSyncDriver], Any]", *args: Any, **kwargs: Any) -> Any: + """Execute a unit of work inside a transaction, retrying on abort. + + If the driver is already bound to a SpannerTransaction, the callable + is executed directly. Otherwise, work is delegated to the database + transaction retry runner. + + Args: + func: Callback taking this driver or a transaction-bound driver. + *args: Positional arguments passed to the callback. + **kwargs: Keyword arguments passed to the callback. + + Returns: + The return value of the callback. + """ + if isinstance(self.connection, SpannerTransaction): + return func(self, *args, **kwargs) + + database = self._get_database() + if database is None: + msg = "run_in_transaction requires an active database or SpannerTransaction context." + raise SQLConversionError(msg) + + def _callback(spanner_transaction: Any, *cb_args: Any, **cb_kwargs: Any) -> Any: + driver = SpannerSyncDriver( + connection=spanner_transaction, + statement_config=self.statement_config, + driver_features=self.driver_features, + ) + driver._config = self._config + call_args = cb_args or args + call_kwargs = cb_kwargs or kwargs + try: + return func(driver, *call_args, **call_kwargs) + except Exception as exc: + cause = getattr(exc, "__cause__", None) + from google.api_core import exceptions as api_exceptions + + if isinstance(exc, api_exceptions.Aborted): + raise + if cause is not None and isinstance(cause, api_exceptions.Aborted): + raise cause from exc + raise + + return database.run_in_transaction(_callback, *args, **kwargs) + + def execute_partitioned_dml( + self, + statement: "SQL | str", + *parameters: Any, + query_options: Any = None, + request_options: Any = None, + exclude_txn_from_change_streams: bool = False, + **kwargs: Any, + ) -> int: + """Execute a Partitioned DML statement across database partitions. + + Args: + statement: The SQL string or SQL object to execute. + *parameters: Positional parameters or parameter mapping. + query_options: Optional Spanner QueryOptions. + request_options: Optional Spanner RequestOptions. + exclude_txn_from_change_streams: Whether to exclude the transaction from change streams. + **kwargs: Additional keyword arguments or parameters. + + Returns: + The number of affected rows. + """ + database = self._get_database() + if database is None: + msg = "Could not resolve Spanner database for partitioned DML execution." + raise SQLConversionError(msg) + + if isinstance(statement, SQL): + sql_statement = statement + else: + sql_statement = self.prepare_statement( + statement, parameters, statement_config=self.statement_config, kwargs=kwargs or None + ) + + sql, raw_params = self._compiled_sql(sql_statement, self.statement_config) + params = raw_params if isinstance(raw_params, dict) else None + coerced_params = self._coerce_params(params) + param_types = self._infer_param_types(params) + + effective_request_options = request_options or self.driver_features.get("request_options") + effective_query_options = query_options or self.driver_features.get("query_options") + + call_kwargs: dict[str, Any] = { + "params": coerced_params, + "param_types": param_types, + "exclude_txn_from_change_streams": exclude_txn_from_change_streams, + } + if effective_query_options is not None: + call_kwargs["query_options"] = effective_query_options + if effective_request_options is not None: + call_kwargs["request_options"] = effective_request_options + + return database.execute_partitioned_dml(sql, **call_kwargs) + def execute( self, statement: "SQL | Statement | QueryBuilder", @@ -439,24 +556,18 @@ def load_from_arrow( arrow_table = self._coerce_arrow_table(source) if overwrite: - delete_sql = f"DELETE FROM {table} WHERE TRUE" - if isinstance(self.connection, SpannerTransaction): - writer = cast("_SpannerWriteProtocol", self.connection) - writer.execute_update(delete_sql) - else: - msg = "Delete requires a Transaction context." - raise SQLConversionError(msg) + self.execute_partitioned_dml(f"DELETE FROM {table} WHERE TRUE") columns, records = self._arrow_table_to_rows(arrow_table) if records: - conn = self.connection - if not isinstance(conn, SpannerTransaction): - msg = "Arrow import requires a Transaction context." - raise SQLConversionError(msg) chunks = self._chunk_mutation_rows(columns, records) if self.driver_features.get("enable_batch_write_api") and not overwrite: self._batch_write_mutations(table, columns, chunks) else: + conn = self.connection + if not isinstance(conn, SpannerTransaction): + msg = "Arrow import requires a Transaction context." + raise SQLConversionError(msg) writer = cast("_SpannerWriteProtocol", conn) for chunk in chunks: writer.insert_or_update(table, columns, chunk) @@ -512,36 +623,57 @@ def resolve_rowcount(self, cursor: "SpannerConnection") -> int: """ return 0 - def _execute_kwargs(self, *, for_read: bool = False) -> dict[str, Any]: + def _execute_kwargs(self, *, for_read: bool = False, for_batch: bool = False) -> dict[str, Any]: kwargs: dict[str, Any] = { key: self.driver_features[key] for key in ("retry", "timeout") if key in self.driver_features } request_options = self.driver_features.get("request_options") if request_options is not None: kwargs["request_options"] = request_options - directed_read_options = self.driver_features.get("directed_read_options") - if for_read and directed_read_options is not None: - kwargs["directed_read_options"] = directed_read_options + if not for_batch: + query_options = self.driver_features.get("query_options") + if query_options is not None: + kwargs["query_options"] = query_options + if for_read and not for_batch: + directed_read_options = self.driver_features.get("directed_read_options") + if directed_read_options is not None: + kwargs["directed_read_options"] = directed_read_options pending = self._pending_execute_options if pending is not None: if pending.request_options is not None: kwargs["request_options"] = pending.request_options + if not for_batch and pending.query_options is not None: + kwargs["query_options"] = pending.query_options if pending.retry is not None: kwargs["retry"] = pending.retry if pending.timeout is not None: kwargs["timeout"] = pending.timeout - if for_read and pending.directed_read_options is not None: + if for_read and not for_batch and pending.directed_read_options is not None: kwargs["directed_read_options"] = pending.directed_read_options + if not for_read and pending.last_statement: + kwargs["last_statement"] = True return kwargs def _pop_execute_options(self, kwargs: dict[str, Any]) -> "_PerCallExecuteOptions | None": - if not any(key in kwargs for key in ("request_options", "directed_read_options", "retry", "timeout")): + if not any( + key in kwargs + for key in ( + "request_options", + "query_options", + "directed_read_options", + "retry", + "timeout", + "last_statement", + ) + ): return None return _PerCallExecuteOptions( request_options=kwargs.pop("request_options", None), + query_options=kwargs.pop("query_options", None), directed_read_options=kwargs.pop("directed_read_options", None), retry=kwargs.pop("retry", None), timeout=kwargs.pop("timeout", None), + last_statement=bool(kwargs.pop("last_statement", False)), ) def _chunk_mutation_rows(self, columns: "list[str]", records: "list[tuple[Any, ...]]") -> "list[list[list[Any]]]": @@ -569,10 +701,9 @@ def _chunk_mutation_rows(self, columns: "list[str]", records: "list[tuple[Any, . def _batch_write_mutations(self, table: str, columns: "list[str]", chunks: "list[list[list[Any]]]") -> None: """High-throughput ingest via the Spanner Batch Write API (one mutation group per chunk).""" - session = cast("object", getattr(self.connection, "_session", None)) - database = cast("Any", getattr(session, "_database", None)) if session is not None else None + database = self._get_database() if database is None: - msg = "Spanner Batch Write API requires a database-backed session." + msg = "Spanner Batch Write API requires a database-backed session or config." raise SQLConversionError(msg) with database.mutation_groups() as mutation_groups: for chunk in chunks: @@ -647,20 +778,24 @@ def rollback(self) -> None: ... class _PerCallExecuteOptions: """Per-call Spanner execution options captured for a single dispatch.""" - __slots__ = ("directed_read_options", "request_options", "retry", "timeout") + __slots__ = ("directed_read_options", "last_statement", "query_options", "request_options", "retry", "timeout") def __init__( self, *, request_options: "RequestOptions | dict[str, Any] | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, + last_statement: bool = False, ) -> None: self.request_options = request_options + self.query_options = query_options self.directed_read_options = directed_read_options self.retry = retry self.timeout = timeout + self.last_statement = last_statement class _SpannerSelectStreamSource: diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index e81ecb153..31db32c85 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -17,6 +17,7 @@ from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction from sqlspec.adapters.spanner.config import SpannerSyncConfig + from sqlspec.adapters.spanner.driver import SpannerSyncDriver class _DatabaseProtocol(Protocol): def run_in_transaction(self, func: "Callable[[Transaction], Any]") -> Any: ... @@ -163,8 +164,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if self._shard_count > 1: update_sql = f"{update_sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = self._build_params(key, new_expires) - types = self._get_param_types(expires_at=True) - self._database().run_in_transaction(_SpannerExecuteUpdateJob(update_sql, params, types)) + self._config.run_in_transaction(lambda driver: driver.execute(update_sql, params)) return spanner_to_bytes(data) @@ -172,7 +172,6 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) params = self._build_params(key, expires_at, data) - types = self._get_param_types(session_id=True, expires_at=True, data=True) update_sql = f""" UPDATE {self._table_name} @@ -187,19 +186,24 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No INSERT {self._table_name} (session_id, data, expires_at, created_at, updated_at) VALUES (@session_id, @data, @expires_at, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()) """ - self._database().run_in_transaction(_SpannerUpsertJob(update_sql, insert_sql, params, types)) + + def _job(driver: "SpannerSyncDriver") -> None: + result = driver.execute(update_sql, params) + if not getattr(result, "rowcount", None): + driver.execute(insert_sql, params) + + self._config.run_in_transaction(_job) def _delete(self, key: str) -> None: sql = f"DELETE FROM {self._table_name} WHERE session_id = @session_id" if self._shard_count > 1: sql = f"{sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = {"session_id": key} - types = self._get_param_types(session_id=True) - self._database().run_in_transaction(_SpannerExecuteUpdateJob(sql, params, types)) + self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name} WHERE TRUE" - self._database().run_in_transaction(_SpannerExecuteUpdateJob(sql)) + self._config.run_in_transaction(lambda driver: driver.execute(sql)) def _exists(self, key: str) -> bool: sql = f""" @@ -234,8 +238,8 @@ def _delete_expired(self) -> int: DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= CURRENT_TIMESTAMP() """ - result = self._database().run_in_transaction(_SpannerExecuteUpdateCountJob(sql)) - return cast("int", result) + result = self._config.run_in_transaction(lambda driver: driver.execute(sql)) + return cast("int", getattr(result, "rowcount", 0)) def _create_table(self) -> None: database = self._config.get_database() @@ -275,43 +279,3 @@ def _index_ddl(self) -> str: def _drop_table_sql(self) -> "list[str]": return [f"DROP INDEX idx_{self._table_name}_expires_at", f"DROP TABLE {self._table_name}"] - - -class _SpannerExecuteUpdateJob: - __slots__ = ("_params", "_sql", "_types") - - def __init__(self, sql: str, params: "dict[str, Any] | None" = None, types: "dict[str, Any] | None" = None) -> None: - self._sql = sql - self._params = params - self._types = types - - def __call__(self, transaction: "Transaction") -> None: - if self._params is None and self._types is None: - transaction.execute_update(self._sql) # type: ignore[no-untyped-call] - return - transaction.execute_update(self._sql, params=self._params or {}, param_types=self._types) # type: ignore[no-untyped-call] - - -class _SpannerUpsertJob: - __slots__ = ("_insert_sql", "_params", "_types", "_update_sql") - - def __init__(self, update_sql: str, insert_sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: - self._update_sql = update_sql - self._insert_sql = insert_sql - self._params = params - self._types = types - - def __call__(self, transaction: "Transaction") -> None: - row_ct = transaction.execute_update(self._update_sql, params=self._params, param_types=self._types) # type: ignore[no-untyped-call] - if row_ct == 0: - transaction.execute_update(self._insert_sql, params=self._params, param_types=self._types) # type: ignore[no-untyped-call] - - -class _SpannerExecuteUpdateCountJob: - __slots__ = ("_sql",) - - def __init__(self, sql: str) -> None: - self._sql = sql - - def __call__(self, transaction: "Transaction") -> int: - return int(transaction.execute_update(self._sql)) # type: ignore[no-untyped-call] diff --git a/sqlspec/adapters/spanner/type_converter.py b/sqlspec/adapters/spanner/type_converter.py index 29907d474..207cb85d0 100644 --- a/sqlspec/adapters/spanner/type_converter.py +++ b/sqlspec/adapters/spanner/type_converter.py @@ -17,7 +17,7 @@ """ import base64 -from datetime import date, datetime, timezone +from datetime import date, datetime, timedelta, timezone from decimal import Decimal from typing import TYPE_CHECKING, Any, cast from uuid import UUID @@ -33,6 +33,8 @@ from sqlspec.protocols import SpannerParamTypesProtocol __all__ = ( + "SPANNER_FLOAT32", + "SPANNER_VECTOR", "bytes_to_spanner", "coerce_params_for_spanner", "infer_spanner_param_types", @@ -42,6 +44,9 @@ "uuid_to_spanner", ) +SPANNER_FLOAT32: str = "FLOAT32" +SPANNER_VECTOR: str = "ARRAY" + _UUID_TYPES: "tuple[type[Any], ...]" = (UUID,) _uuid_utils_uuid = import_optional_attr("uuid_utils", "UUID") if _uuid_utils_uuid is not None: @@ -168,10 +173,18 @@ def coerce_params_for_spanner( coerced: dict[str, Any] = {} changed = False for key, value in params.items(): + declared_type = None if type(value) is TypedParameter: + declared_type = value.original_type value = value.value changed = True - if isinstance(value, _UUID_TYPES): + if declared_type in ("ARRAY", "VECTOR", "vector") and isinstance(value, (list, tuple)): + coerced[key] = [float(x) for x in value] + changed = True + elif declared_type in ("FLOAT32", "float32") and value is not None: + coerced[key] = float(value) + changed = True + elif isinstance(value, _UUID_TYPES): if enable_uuid_conversion: coerced[key] = str(value) changed = True @@ -202,7 +215,7 @@ def coerce_params_for_spanner( return coerced if changed else params -_NULL_PARAM_TYPE_NAMES: "dict[type[Any], str]" = { +_NULL_PARAM_TYPE_NAMES: "dict[type[Any] | str, str]" = { bool: "BOOL", int: "INT64", float: "FLOAT64", @@ -211,10 +224,51 @@ def coerce_params_for_spanner( datetime: "TIMESTAMP", date: "DATE", Decimal: "NUMERIC", + timedelta: "INTERVAL", UUID: "STRING", + "FLOAT32": "FLOAT32", + "float32": "FLOAT32", + "FLOAT64": "FLOAT64", + "float64": "FLOAT64", + "INTERVAL": "INTERVAL", + "interval": "INTERVAL", } +def _infer_sequence_param_type(value: Any, param_types: Any, json_type: Any) -> Any | None: + """Infer Spanner parameter type for sequence values. + + Args: + value: Sequence value to inspect. + param_types: The Spanner param_types module. + json_type: The Spanner JSON param type. + + Returns: + Spanner Array type, JSON type, or None if sequence is empty or unhandled. + """ + if should_json_encode_sequence(value): + return json_type + sequence = list(value) + if not sequence: + return None + first = sequence[0] + if isinstance(first, int): + return param_types.Array(param_types.INT64) + if isinstance(first, str): + return param_types.Array(param_types.STRING) + if isinstance(first, float): + return param_types.Array(param_types.FLOAT64) + if isinstance(first, bool): + return param_types.Array(param_types.BOOL) + if isinstance(first, Decimal): + return param_types.Array(param_types.NUMERIC) + if isinstance(first, timedelta): + interval_type = getattr(param_types, "INTERVAL", None) + if interval_type is not None: + return param_types.Array(interval_type) + return None + + def infer_spanner_param_types(params: "dict[str, Any] | None") -> "dict[str, Any]": """Infer Spanner param_types from Python values. @@ -232,17 +286,39 @@ def infer_spanner_param_types(params: "dict[str, Any] | None") -> "dict[str, Any types: dict[str, Any] = {} json_type = _json_param_type() for key, raw_value in params.items(): - value = raw_value.value if type(raw_value) is TypedParameter else raw_value + is_typed = type(raw_value) is TypedParameter + value = raw_value.value if is_typed else raw_value + declared = raw_value.original_type if is_typed else None if value is None: null_type = _null_param_type(raw_value, param_types) if null_type is not None: types[key] = null_type - elif isinstance(value, bool): + continue + if declared in ("FLOAT32", "float32"): + types[key] = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) + continue + if declared in ("ARRAY", "VECTOR", "vector"): + float32_type = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) + types[key] = param_types.Array(float32_type) + continue + if declared in ("FLOAT64", "float64"): + types[key] = param_types.FLOAT64 + continue + if declared == "ARRAY": + types[key] = param_types.Array(param_types.FLOAT64) + continue + if isinstance(value, bool): types[key] = param_types.BOOL elif isinstance(value, int): types[key] = param_types.INT64 elif isinstance(value, float): types[key] = param_types.FLOAT64 + elif isinstance(value, Decimal): + types[key] = param_types.NUMERIC + elif isinstance(value, timedelta): + interval_type = getattr(param_types, "INTERVAL", None) + if interval_type is not None: + types[key] = interval_type elif isinstance(value, _STRING_PARAM_TYPES): types[key] = param_types.STRING elif isinstance(value, bytes): @@ -254,21 +330,9 @@ def infer_spanner_param_types(params: "dict[str, Any] | None") -> "dict[str, Any elif isinstance(value, (dict, json_object_type)): types[key] = json_type elif isinstance(value, (list, tuple)): - if should_json_encode_sequence(value): - types[key] = json_type - continue - sequence = list(value) - if not sequence: - continue - first = sequence[0] - if isinstance(first, int): - types[key] = param_types.Array(param_types.INT64) - elif isinstance(first, str): - types[key] = param_types.Array(param_types.STRING) - elif isinstance(first, float): - types[key] = param_types.Array(param_types.FLOAT64) - elif isinstance(first, bool): - types[key] = param_types.Array(param_types.BOOL) + seq_type = _infer_sequence_param_type(value, param_types, json_type) + if seq_type is not None: + types[key] = seq_type return types @@ -290,7 +354,16 @@ def _null_param_type(raw_value: Any, param_types: "SpannerParamTypesProtocol") - declared = raw_value.original_type if type(raw_value) is TypedParameter else None if declared is None: return None + if declared in ("ARRAY", "VECTOR", "vector"): + float32_type = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) + return param_types.Array(float32_type) + if declared == "ARRAY": + return param_types.Array(param_types.FLOAT64) resolver = _NULL_PARAM_TYPE_NAMES.get(declared) + if resolver == "FLOAT32": + return getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) + if resolver == "INTERVAL": + return getattr(param_types, "INTERVAL", None) return getattr(param_types, resolver) if resolver is not None else None diff --git a/sqlspec/core/parameters/_types.py b/sqlspec/core/parameters/_types.py index 703cd9fb5..e7fd25b54 100644 --- a/sqlspec/core/parameters/_types.py +++ b/sqlspec/core/parameters/_types.py @@ -176,7 +176,9 @@ class TypedParameter: __slots__ = TYPED_PARAMETER_SLOTS - def __init__(self, value: Any, original_type: "type | None" = None, semantic_name: "str | None" = None) -> None: + def __init__( + self, value: Any, original_type: "type | str | None" = None, semantic_name: "str | None" = None + ) -> None: self.value = value self.original_type = original_type or type(value) self.semantic_name = semantic_name @@ -199,7 +201,8 @@ def __eq__(self, other: object) -> bool: def __repr__(self) -> str: name_part = f", semantic_name='{self.semantic_name}'" if self.semantic_name else "" - return f"TypedParameter({self.value!r}, original_type={self.original_type.__name__}{name_part})" + type_name = getattr(self.original_type, "__name__", str(self.original_type)) + return f"TypedParameter({self.value!r}, original_type={type_name}{name_part})" def __reduce__(self) -> "tuple[Any, ...]": """Reconstruct via ``TypedParameter(value, original_type, semantic_name)``.""" diff --git a/sqlspec/protocols.py b/sqlspec/protocols.py index d6b25e75e..9d85d8d5a 100644 --- a/sqlspec/protocols.py +++ b/sqlspec/protocols.py @@ -282,7 +282,10 @@ class SpannerParamTypesProtocol(SupportsJsonTypeProtocol, Protocol): BOOL: Any INT64: Any + FLOAT32: Any FLOAT64: Any + NUMERIC: Any + INTERVAL: Any STRING: Any BYTES: Any TIMESTAMP: Any diff --git a/tests/unit/adapters/test_spanner/test_batch_write_api.py b/tests/unit/adapters/test_spanner/test_batch_write_api.py index 8a002e02a..30df6b68f 100644 --- a/tests/unit/adapters/test_spanner/test_batch_write_api.py +++ b/tests/unit/adapters/test_spanner/test_batch_write_api.py @@ -53,10 +53,15 @@ def batch_write(self, request_options: Any = None, exclude_txn_from_change_strea class _FakeDatabase: def __init__(self) -> None: self.mutation_groups_obj = _FakeMutationGroups() + self.partitioned_dml_calls: list[str] = [] def mutation_groups(self) -> _FakeMutationGroups: return self.mutation_groups_obj + def execute_partitioned_dml(self, dml: str, **kwargs: Any) -> int: + self.partitioned_dml_calls.append(dml) + return 0 + class _FakeSession: def __init__(self, database: _FakeDatabase) -> None: @@ -114,6 +119,8 @@ def test_batch_write_overwrite_uses_transactional_mutations(batch_write_driver: conn = cast("_FakeBatchTransaction", batch_write_driver.connection) batch_write_driver.load_from_arrow("users", pa.table({"id": [1]}), overwrite=True) - assert conn.execute_update_calls and "DELETE FROM users WHERE TRUE" in conn.execute_update_calls[0] + assert ( + conn.database.partitioned_dml_calls and "DELETE FROM users WHERE TRUE" in conn.database.partitioned_dml_calls[0] + ) assert conn.insert_or_update_calls == [("users", ["id"], [[1]])] assert conn.database.mutation_groups_obj.batch_write_calls == 0 diff --git a/tests/unit/adapters/test_spanner/test_litestar_store.py b/tests/unit/adapters/test_spanner/test_litestar_store.py index 6022cc66c..8f4d7525e 100644 --- a/tests/unit/adapters/test_spanner/test_litestar_store.py +++ b/tests/unit/adapters/test_spanner/test_litestar_store.py @@ -12,59 +12,47 @@ def _mock_database() -> MagicMock: def test_set_uses_run_in_transaction() -> None: - """Verify _set uses database.run_in_transaction for write operations.""" - mock_db = _mock_database() - + """Verify _set uses config.run_in_transaction for write operations.""" config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.get_database.return_value = mock_db store = SpannerSyncStore(config) store._set("s1", b"data", None) # pyright: ignore - mock_db.run_in_transaction.assert_called_once() + config.run_in_transaction.assert_called_once() def test_delete_uses_run_in_transaction() -> None: - """Verify _delete uses database.run_in_transaction for write operations.""" - mock_db = _mock_database() - + """Verify _delete uses config.run_in_transaction for write operations.""" config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.get_database.return_value = mock_db store = SpannerSyncStore(config) store._delete("s1") # pyright: ignore - mock_db.run_in_transaction.assert_called_once() + config.run_in_transaction.assert_called_once() def test_delete_all_uses_run_in_transaction() -> None: - """Verify _delete_all uses database.run_in_transaction for write operations.""" - mock_db = _mock_database() - + """Verify _delete_all uses config.run_in_transaction for write operations.""" config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.get_database.return_value = mock_db store = SpannerSyncStore(config) store._delete_all() # pyright: ignore - mock_db.run_in_transaction.assert_called_once() + config.run_in_transaction.assert_called_once() def test_delete_expired_uses_run_in_transaction() -> None: - """Verify _delete_expired uses database.run_in_transaction for write operations.""" - mock_db = _mock_database() - + """Verify _delete_expired uses config.run_in_transaction for write operations.""" config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.get_database.return_value = mock_db store = SpannerSyncStore(config) store._delete_expired() # pyright: ignore - mock_db.run_in_transaction.assert_called_once() + config.run_in_transaction.assert_called_once() def _context_manager_yielding(value: Any) -> Any: diff --git a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py index f459d5848..fbd76fb90 100644 --- a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py +++ b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py @@ -22,6 +22,10 @@ def __init__(self) -> None: self.insert_or_update_calls: list[tuple[str, list[str], list[list[Any]]]] = [] self.execute_update_calls: list[str] = [] self.committed = None + from unittest.mock import MagicMock + + self._database = MagicMock() + self._database.execute_partitioned_dml.return_value = 0 def insert_or_update(self, table: str, columns: Any, values: Any) -> None: self.insert_or_update_calls.append((table, list(columns), [list(v) for v in values])) @@ -70,8 +74,9 @@ def test_load_from_arrow_overwrite_deletes_then_mutates(mutations_driver: Spanne mutations_driver.load_from_arrow("users", arrow_table, overwrite=True) - assert txn.execute_update_calls - assert "DELETE FROM users WHERE TRUE" in txn.execute_update_calls[0] + txn._database.execute_partitioned_dml.assert_called_once() + sql = txn._database.execute_partitioned_dml.call_args[0][0] + assert "DELETE FROM users WHERE TRUE" in sql assert len(txn.insert_or_update_calls) == 1 diff --git a/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py b/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py new file mode 100644 index 000000000..1c45d95d3 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py @@ -0,0 +1,35 @@ +"""Unit tests for load_from_arrow(overwrite=True) using Partitioned DML.""" + +from unittest.mock import MagicMock, patch + +import pyarrow as pa + +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + +CAPABILITIES = { + "arrow_export_enabled": True, + "arrow_import_enabled": True, + "parquet_export_enabled": True, + "parquet_import_enabled": True, + "partition_strategies": ["fixed"], +} + + +def test_load_from_arrow_overwrite_uses_partitioned_dml() -> None: + """Verify load_from_arrow(overwrite=True) calls execute_partitioned_dml for truncation.""" + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 1000 + + mock_connection = MagicMock() + mock_connection._session._database = mock_db + + driver = SpannerSyncDriver(connection=mock_connection, driver_features={"storage_capabilities": CAPABILITIES}) + + arrow_table = pa.table({"id": [1, 2], "name": ["a", "b"]}) + + with patch.object(SpannerSyncDriver, "_arrow_table_to_rows", return_value=(["id", "name"], [])): + driver.load_from_arrow("users", arrow_table, overwrite=True) + + mock_db.execute_partitioned_dml.assert_called_once() + sql = mock_db.execute_partitioned_dml.call_args[0][0] + assert "DELETE FROM users WHERE TRUE" in sql diff --git a/tests/unit/adapters/test_spanner/test_spanner_batch_write.py b/tests/unit/adapters/test_spanner/test_spanner_batch_write.py new file mode 100644 index 000000000..6e41eaf66 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_batch_write.py @@ -0,0 +1,59 @@ +"""Unit tests for load_from_arrow with enable_batch_write_api.""" + +from unittest.mock import MagicMock + +import pyarrow as pa +import pytest + +from sqlspec.adapters.spanner.driver import SpannerSyncDriver +from sqlspec.exceptions import SQLConversionError + +CAPABILITIES = { + "arrow_export_enabled": True, + "arrow_import_enabled": True, + "parquet_export_enabled": True, + "parquet_import_enabled": True, + "partition_strategies": ["fixed"], +} + + +def test_batch_write_succeeds_without_transaction() -> None: + """Verify load_from_arrow with enable_batch_write_api succeeds on non-transaction connection.""" + mock_db = MagicMock() + mock_mg = MagicMock() + mock_db.mutation_groups.return_value.__enter__.return_value = mock_mg + mock_group = MagicMock() + mock_mg.group.return_value = mock_group + mock_response = MagicMock() + mock_response.status = None + mock_mg.batch_write.return_value = [mock_response] + + mock_snapshot = MagicMock() + mock_snapshot._session._database = mock_db + + driver = SpannerSyncDriver( + connection=mock_snapshot, driver_features={"storage_capabilities": CAPABILITIES, "enable_batch_write_api": True} + ) + + arrow_table = pa.table({"id": [1, 2], "name": ["alice", "bob"]}) + + job = driver.load_from_arrow("users", arrow_table) + assert job.telemetry["rows_processed"] == 2 + mock_db.mutation_groups.assert_called_once() + mock_mg.batch_write.assert_called_once() + mock_group.insert_or_update.assert_called_once() + + +def test_standard_insert_requires_transaction() -> None: + """Verify load_from_arrow without enable_batch_write_api still requires a SpannerTransaction.""" + mock_snapshot = MagicMock() + + driver = SpannerSyncDriver( + connection=mock_snapshot, + driver_features={"storage_capabilities": CAPABILITIES, "enable_batch_write_api": False}, + ) + + arrow_table = pa.table({"id": [1, 2], "name": ["alice", "bob"]}) + + with pytest.raises(SQLConversionError, match=r"Arrow import requires a Transaction context\."): + driver.load_from_arrow("users", arrow_table) diff --git a/tests/unit/adapters/test_spanner/test_spanner_json.py b/tests/unit/adapters/test_spanner/test_spanner_json.py new file mode 100644 index 000000000..824a80d3d --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_json.py @@ -0,0 +1,54 @@ +"""Unit tests for Spanner JsonObject direct unwrapping optimization.""" + +from unittest.mock import MagicMock + +from google.cloud.spanner_v1.data_types import JsonObject + +from sqlspec.adapters.spanner.core import _convert_json_row_value +from sqlspec.utils.serializers import from_json + + +class MonitoredJsonObject(JsonObject): + """JsonObject subclass that tracks calls to serialize().""" + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + self.serialize_called = False + + def serialize(self) -> str | None: + self.serialize_called = True + return super().serialize() + + +def test_convert_json_row_value_unwraps_without_serialize() -> None: + """Verify that default deserializer unwraps JsonObject directly without calling serialize().""" + obj = MonitoredJsonObject({"key": "val", "nested": [1, 2]}) + res = _convert_json_row_value(obj, json_deserializer=from_json) + assert res == {"key": "val", "nested": [1, 2]} + assert not obj.serialize_called + + +def test_convert_json_row_value_calls_serialize_for_custom_deserializer() -> None: + """Verify that a custom string deserializer invokes serialize().""" + obj = MonitoredJsonObject({"key": "val"}) + custom_deserializer = MagicMock(return_value={"custom": True}) + res = _convert_json_row_value(obj, json_deserializer=custom_deserializer) + assert res == {"custom": True} + assert obj.serialize_called + custom_deserializer.assert_called_once_with('{"key":"val"}') + + +def test_convert_json_row_value_null_json() -> None: + """Verify that null JsonObject unwraps directly to None.""" + obj = MonitoredJsonObject(None) + res = _convert_json_row_value(obj, json_deserializer=from_json) + assert res is None + assert not obj.serialize_called + + +def test_convert_json_row_value_array_json() -> None: + """Verify that array JsonObject unwraps directly to a list.""" + obj = MonitoredJsonObject([1, 2, 3]) + res = _convert_json_row_value(obj, json_deserializer=from_json) + assert res == [1, 2, 3] + assert not obj.serialize_called diff --git a/tests/unit/adapters/test_spanner/test_spanner_last_statement.py b/tests/unit/adapters/test_spanner/test_spanner_last_statement.py new file mode 100644 index 000000000..c0dc62d0a --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_last_statement.py @@ -0,0 +1,58 @@ +"""Unit tests for Spanner last_statement execution and commit behavior.""" + +from datetime import datetime, timezone +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.config import SpannerConnectionContext, SpannerSyncConfig +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_execute_passes_last_statement_to_writer() -> None: + """Verify that driver.execute forwards last_statement=True to writer.execute_update.""" + mock_cursor = MagicMock() + mock_cursor.execute_update.return_value = 1 + mock_cursor.committed = None + + driver = SpannerSyncDriver(connection=mock_cursor) + driver.execute("UPDATE t SET x = 1 WHERE id = 'a'", last_statement=True) + + mock_cursor.execute_update.assert_called_once() + _, kwargs = mock_cursor.execute_update.call_args + assert kwargs.get("last_statement") is True + + +def test_driver_commit_noop_when_transaction_already_committed() -> None: + """Verify that driver.commit is a no-op when writer.committed is set.""" + mock_cursor = MagicMock() + mock_cursor.committed = datetime.now(timezone.utc) + mock_cursor.commit = MagicMock() + + driver = SpannerSyncDriver(connection=mock_cursor) + driver.commit() + + mock_cursor.commit.assert_not_called() + + +def test_connection_context_exit_noop_when_already_committed() -> None: + """Verify that SpannerConnectionContext.__exit__ skips commit when txn.committed is set.""" + mock_txn = MagicMock() + mock_txn._transaction_id = b"tx1" + mock_txn.committed = datetime.now(timezone.utc) + mock_txn.commit = MagicMock() + + mock_session = MagicMock() + mock_session.transaction.return_value = mock_txn + + mock_db = MagicMock() + mock_db.sessions_manager.put_session = MagicMock() + + config = MagicMock(spec=SpannerSyncConfig) + config.get_database.return_value = mock_db + + ctx = SpannerConnectionContext(config, transaction=True) + ctx._session = mock_session + ctx._connection = mock_txn + + ctx.__exit__(None, None, None) + + mock_txn.commit.assert_not_called() diff --git a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py new file mode 100644 index 000000000..cd42e46c9 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py @@ -0,0 +1,146 @@ +"""Unit tests for Spanner Partitioned DML execution on driver and config.""" + +from unittest.mock import MagicMock, patch + +from google.cloud.spanner_v1.types.type import TypeCode + +from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.core import default_statement_config +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_driver_execute_partitioned_dml() -> None: + """Verify that driver.execute_partitioned_dml delegates to database.execute_partitioned_dml.""" + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 42 + + mock_connection = MagicMock() + mock_connection._session._database = mock_db + + driver = SpannerSyncDriver( + connection=mock_connection, statement_config=default_statement_config, driver_features={} + ) + + rows = driver.execute_partitioned_dml("DELETE FROM large_table WHERE active = FALSE") + assert rows == 42 + mock_db.execute_partitioned_dml.assert_called_once() + sql = mock_db.execute_partitioned_dml.call_args[0][0] + assert "DELETE FROM large_table WHERE active = FALSE" in sql + + +def test_driver_execute_partitioned_dml_with_parameters() -> None: + """Verify parameters and types are coerced and passed to execute_partitioned_dml.""" + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 10 + + mock_connection = MagicMock() + mock_connection._session._database = mock_db + + driver = SpannerSyncDriver( + connection=mock_connection, statement_config=default_statement_config, driver_features={} + ) + + rows = driver.execute_partitioned_dml( + "UPDATE large_table SET status = :status WHERE threshold > :limit", {"status": "archived", "limit": 100} + ) + assert rows == 10 + mock_db.execute_partitioned_dml.assert_called_once() + _, kwargs = mock_db.execute_partitioned_dml.call_args + assert kwargs["params"] == {"status": "archived", "limit": 100} + assert "status" in kwargs["param_types"] + assert kwargs["param_types"]["status"].code == TypeCode.STRING + assert "limit" in kwargs["param_types"] + assert kwargs["param_types"]["limit"].code == TypeCode.INT64 + + +def test_config_execute_partitioned_dml() -> None: + """Verify that config.execute_partitioned_dml delegates to get_database().execute_partitioned_dml.""" + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 100 + + with patch.object(config, "get_database", return_value=mock_db): + rows = config.execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") + + assert rows == 100 + mock_db.execute_partitioned_dml.assert_called_once() + sql = mock_db.execute_partitioned_dml.call_args[0][0] + assert "DELETE FROM large_table WHERE TRUE" in sql + + +def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: + """Verify executing partitioned DML with SQL object, query options, and request options.""" + from sqlspec.core import SQL + + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 50 + + mock_connection = MagicMock() + mock_connection._session._database = mock_db + + driver = SpannerSyncDriver( + connection=mock_connection, statement_config=default_statement_config, driver_features={} + ) + + statement = SQL("DELETE FROM large_table WHERE expired = TRUE", statement_config=default_statement_config) + mock_query_options = MagicMock() + mock_request_options = MagicMock() + + rows = driver.execute_partitioned_dml( + statement, + query_options=mock_query_options, + request_options=mock_request_options, + exclude_txn_from_change_streams=True, + ) + assert rows == 50 + mock_db.execute_partitioned_dml.assert_called_once() + _, kwargs = mock_db.execute_partitioned_dml.call_args + assert kwargs["query_options"] is mock_query_options + assert kwargs["request_options"] is mock_request_options + assert kwargs["exclude_txn_from_change_streams"] is True + + +def test_driver_execute_partitioned_dml_no_database_raises() -> None: + """Verify error raised when database cannot be resolved.""" + import pytest + + from sqlspec.exceptions import SQLConversionError + + mock_connection = MagicMock() + mock_connection._session = None + mock_connection._database = None + + driver = SpannerSyncDriver( + connection=mock_connection, statement_config=default_statement_config, driver_features={} + ) + + with pytest.raises(SQLConversionError, match="Could not resolve Spanner database"): + driver.execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") + + +def test_config_execute_partitioned_dml_with_parameters_and_options() -> None: + """Verify config partitioned DML forwards parameters, types, and options.""" + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + mock_db = MagicMock() + mock_db.execute_partitioned_dml.return_value = 25 + + mock_query_options = MagicMock() + mock_request_options = MagicMock() + + with patch.object(config, "get_database", return_value=mock_db): + rows = config.execute_partitioned_dml( + "UPDATE items SET status = :status WHERE id = :id", + {"status": "deleted", "id": 5}, + query_options=mock_query_options, + request_options=mock_request_options, + exclude_txn_from_change_streams=True, + ) + + assert rows == 25 + mock_db.execute_partitioned_dml.assert_called_once() + _, kwargs = mock_db.execute_partitioned_dml.call_args + assert kwargs["params"] == {"status": "deleted", "id": 5} + assert "status" in kwargs["param_types"] + assert kwargs["query_options"] is mock_query_options + assert kwargs["request_options"] is mock_request_options + assert kwargs["exclude_txn_from_change_streams"] is True diff --git a/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py b/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py new file mode 100644 index 000000000..2df4e3b7c --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py @@ -0,0 +1,51 @@ +"""Unit tests for Spanner PingingPool fallback and configuration.""" + +from google.cloud.spanner_v1.pool import PingingPool + +from sqlspec.adapters.spanner.config import SpannerSyncConfig + + +def test_disable_multiplexed_sessions_defaults_to_pinging_pool() -> None: + """Verify that disabling multiplexed sessions defaults to PingingPool with 1800s interval.""" + config = SpannerSyncConfig( + connection_config={ + "project": "test-project", + "instance_id": "test-instance", + "database_id": "test-db", + "enable_multiplexed_sessions": False, + } + ) + assert config.connection_config.get("pool_type") is PingingPool + assert config.connection_config.get("ping_interval") == 1800 + + pool = config.provide_pool() + assert isinstance(pool, PingingPool) + assert pool._delta.total_seconds() == 1800 + + +def test_pinging_pool_custom_ping_interval() -> None: + """Verify that custom ping_interval is respected when configuring PingingPool.""" + config = SpannerSyncConfig( + connection_config={ + "project": "test-project", + "instance_id": "test-instance", + "database_id": "test-db", + "enable_multiplexed_sessions": False, + "ping_interval": 900, + } + ) + assert config.connection_config.get("ping_interval") == 900 + + pool = config.provide_pool() + assert isinstance(pool, PingingPool) + assert pool._delta.total_seconds() == 900 + + +def test_provide_pool_fallback_defaults_to_pinging_pool() -> None: + """Verify that calling provide_pool directly falls back to PingingPool with default ping_interval.""" + config = SpannerSyncConfig( + connection_config={"project": "test-project", "instance_id": "test-instance", "database_id": "test-db"} + ) + pool = config.provide_pool() + assert isinstance(pool, PingingPool) + assert pool._delta.total_seconds() == 1800 diff --git a/tests/unit/adapters/test_spanner/test_spanner_pool.py b/tests/unit/adapters/test_spanner/test_spanner_pool.py new file mode 100644 index 000000000..8dfbb52c3 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_pool.py @@ -0,0 +1,71 @@ +"""Unit tests for Spanner session pool configuration and multiplexed pooling.""" + +from unittest.mock import MagicMock, patch + +from google.cloud.spanner_v1.pool import BurstyPool + +from sqlspec.adapters.spanner.config import SpannerSyncConfig + + +def test_multiplexed_session_pool_default() -> None: + """Verify that default SpannerSyncConfig uses multiplexed pooling and omits pool from database args.""" + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + assert config.connection_config.get("pool_type") is None + + mock_db = MagicMock() + mock_instance = MagicMock() + mock_instance.database.return_value = mock_db + mock_client = MagicMock() + mock_client.instance.return_value = mock_instance + + with patch.object(config, "_get_client", return_value=mock_client): + db = config.get_database() + + assert db is mock_db + mock_instance.database.assert_called_once() + _, kwargs = mock_instance.database.call_args + assert "pool" not in kwargs + + +def test_explicit_pool_type_preserved() -> None: + """Verify that explicit pool_type in connection_config is instantiated and forwarded.""" + config = SpannerSyncConfig( + connection_config={"project": "p", "instance_id": "i", "database_id": "d", "pool_type": BurstyPool} + ) + assert config.connection_config.get("pool_type") is BurstyPool + + mock_db = MagicMock() + mock_instance = MagicMock() + mock_instance.database.return_value = mock_db + mock_client = MagicMock() + mock_client.instance.return_value = mock_instance + + with patch.object(config, "_get_client", return_value=mock_client): + db = config.get_database() + + assert db is mock_db + mock_instance.database.assert_called_once() + _, kwargs = mock_instance.database.call_args + assert "pool" in kwargs + assert isinstance(kwargs["pool"], BurstyPool) + + +def test_disable_multiplexed_sessions_uses_legacy_pool() -> None: + """Verify that enable_multiplexed_sessions=False constructs an explicit session pool.""" + config = SpannerSyncConfig( + connection_config={"project": "p", "instance_id": "i", "database_id": "d", "enable_multiplexed_sessions": False} + ) + + mock_db = MagicMock() + mock_instance = MagicMock() + mock_instance.database.return_value = mock_db + mock_client = MagicMock() + mock_client.instance.return_value = mock_instance + + with patch.object(config, "_get_client", return_value=mock_client): + db = config.get_database() + + assert db is mock_db + mock_instance.database.assert_called_once() + _, kwargs = mock_instance.database.call_args + assert "pool" in kwargs diff --git a/tests/unit/adapters/test_spanner/test_spanner_query_options.py b/tests/unit/adapters/test_spanner/test_spanner_query_options.py new file mode 100644 index 000000000..665875388 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_query_options.py @@ -0,0 +1,101 @@ +"""Unit tests for Spanner QueryOptions forwarding.""" + +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.core import default_statement_config +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_driver_execute_with_statement_query_options() -> None: + """Verify driver.execute passes query_options to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1]] + mock_cursor.execute_sql.return_value = mock_result_set + + driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) + + query_opts = {"optimizer_version": "6", "optimizer_statistics_package": "latest"} + driver.execute("SELECT 1", query_options=query_opts) + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts + + +def test_driver_feature_query_options() -> None: + """Verify driver-level query_options are passed to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1]] + mock_cursor.execute_sql.return_value = mock_result_set + + query_opts = {"optimizer_version": "latest"} + driver = SpannerSyncDriver( + connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} + ) + + driver.execute("SELECT 1") + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts + + +def test_config_provide_session_query_options() -> None: + """Verify config.provide_session forwards query_options to driver_features.""" + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + query_opts = {"optimizer_version": "5"} + features = config._session_driver_features( + request_options=None, directed_read_options=None, query_options=query_opts, retry=None, timeout=None + ) + assert features["query_options"] == query_opts + + +def test_driver_execute_many_omits_query_options() -> None: + """Verify execute_many does not forward query_options to batch_update.""" + mock_cursor = MagicMock() + mock_cursor.batch_update.return_value = (None, [1, 1]) + + query_opts = {"optimizer_version": "latest"} + driver = SpannerSyncDriver( + connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} + ) + + driver.execute_many("INSERT INTO t (id) VALUES (:id)", [{"id": 1}, {"id": 2}], query_options=query_opts) + + mock_cursor.batch_update.assert_called_once() + _, kwargs = mock_cursor.batch_update.call_args + assert "query_options" not in kwargs + + +def test_driver_select_stream_query_options() -> None: + """Verify select_stream passes query_options to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1], [2]] + mock_cursor.execute_sql.return_value = mock_result_set + + query_opts = {"optimizer_version": "6"} + driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) + + stream = driver.select_stream("SELECT 1", query_options=query_opts) + rows = list(stream) if stream else [] + assert len(rows) == 2 + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts diff --git a/tests/unit/adapters/test_spanner/test_spanner_request_options.py b/tests/unit/adapters/test_spanner/test_spanner_request_options.py new file mode 100644 index 000000000..1aff60507 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_request_options.py @@ -0,0 +1,105 @@ +"""Unit tests for Spanner RequestOptions and DirectedReadOptions propagation.""" + +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.core import default_statement_config +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_execute_select_forwards_request_and_directed_read_options() -> None: + """Verify execute for SELECT forwards request_options and directed_read_options.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1]] + mock_cursor.execute_sql.return_value = mock_result_set + + driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) + + req_opts = {"request_tag": "select-tag", "priority": 1} + directed_read = MagicMock() + driver.execute("SELECT 1", request_options=req_opts, directed_read_options=directed_read) + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("request_options") == req_opts + assert kwargs.get("directed_read_options") is directed_read + + +def test_execute_update_omits_directed_read_options() -> None: + """Verify execute for UPDATE forwards request_options but strips directed_read_options.""" + mock_cursor = MagicMock() + mock_cursor.execute_update.return_value = 1 + + driver = SpannerSyncDriver( + connection=mock_cursor, + statement_config=default_statement_config, + driver_features={"directed_read_options": MagicMock()}, + ) + + req_opts = {"request_tag": "update-tag", "transaction_tag": "tx-tag"} + per_call_directed = MagicMock() + driver.execute("UPDATE t SET x = 1", request_options=req_opts, directed_read_options=per_call_directed) + + mock_cursor.execute_update.assert_called_once() + _, kwargs = mock_cursor.execute_update.call_args + assert kwargs.get("request_options") == req_opts + assert "directed_read_options" not in kwargs + + +def test_execute_many_forwards_request_options_and_omits_read_options() -> None: + """Verify execute_many forwards request_options and omits directed_read_options and query_options.""" + mock_cursor = MagicMock() + mock_cursor.batch_update.return_value = (None, [1, 1]) + + driver = SpannerSyncDriver( + connection=mock_cursor, + statement_config=default_statement_config, + driver_features={"query_options": {"optimizer_version": "6"}, "directed_read_options": MagicMock()}, + ) + + req_opts = {"request_tag": "batch-tag", "priority": 2} + driver.execute_many( + "INSERT INTO t (id) VALUES (:id)", + [{"id": 1}, {"id": 2}], + request_options=req_opts, + directed_read_options=MagicMock(), + query_options={"optimizer_version": "5"}, + ) + + mock_cursor.batch_update.assert_called_once() + _, kwargs = mock_cursor.batch_update.call_args + assert kwargs.get("request_options") == req_opts + assert "directed_read_options" not in kwargs + assert "query_options" not in kwargs + + +def test_execute_script_forwards_options_appropriately() -> None: + """Verify execute_script separates read and write options across script statements.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_cursor.execute_sql.return_value = mock_result_set + mock_cursor.execute_update.return_value = 1 + + req_opts = {"request_tag": "script-tag"} + directed_read = MagicMock() + driver = SpannerSyncDriver( + connection=mock_cursor, + statement_config=default_statement_config, + driver_features={"request_options": req_opts, "directed_read_options": directed_read}, + ) + + driver.execute_script("SELECT 1; UPDATE t SET x = 1;") + + mock_cursor.execute_sql.assert_called_once() + _, read_kwargs = mock_cursor.execute_sql.call_args + assert read_kwargs.get("request_options") == req_opts + assert read_kwargs.get("directed_read_options") is directed_read + + mock_cursor.execute_update.assert_called_once() + _, write_kwargs = mock_cursor.execute_update.call_args + assert write_kwargs.get("request_options") == req_opts + assert "directed_read_options" not in write_kwargs diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py new file mode 100644 index 000000000..fe29d8021 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -0,0 +1,66 @@ +"""Unit tests for ADK and Litestar store transaction routing.""" + +from typing import Any +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.adk import SpannerSyncADKStore +from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.driver import SpannerSyncDriver +from sqlspec.adapters.spanner.litestar import SpannerSyncStore + + +def test_adk_store_run_write_routes_through_config_run_in_transaction() -> None: + """Verify that SpannerSyncADKStore._run_write executes via config.run_in_transaction.""" + config = MagicMock(spec=SpannerSyncConfig) + executed_statements: list[tuple[str, Any]] = [] + + def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_driver.execute.side_effect = lambda sql, params=None, param_types=None, **kw: executed_statements.append(( + sql, + params, + )) + return func(mock_driver, *args, **kwargs) + + config.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) + + store = SpannerSyncADKStore(config=config) + statements = [ + ("INSERT INTO t (id) VALUES (@id)", {"id": "1"}, {"id": MagicMock()}), + ("INSERT INTO t (id) VALUES (@id)", {"id": "2"}, {"id": MagicMock()}), + ] + store._run_write(statements) + + config.run_in_transaction.assert_called_once() + assert len(executed_statements) == 2 + + +def test_litestar_store_writes_route_through_config_run_in_transaction() -> None: + """Verify that SpannerSyncStore write operations execute via config.run_in_transaction.""" + config = MagicMock(spec=SpannerSyncConfig) + config.extension_config = {"litestar": {"session_table": "sessions"}} + executed_sqls: list[str] = [] + + def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_result = MagicMock() + mock_result.rowcount = 1 + mock_driver.execute.side_effect = lambda sql, *a, **kw: (executed_sqls.append(str(sql)), mock_result)[1] + return func(mock_driver, *args, **kwargs) + + config.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) + + store = SpannerSyncStore(config=config) + + store._set("session_1", b"payload", expires_in=3600) + assert config.run_in_transaction.call_count == 1 + + store._delete("session_1") + assert config.run_in_transaction.call_count == 2 + + store._delete_all() + assert config.run_in_transaction.call_count == 3 + + expired_count = store._delete_expired() + assert config.run_in_transaction.call_count == 4 + assert expired_count == 1 diff --git a/tests/unit/adapters/test_spanner/test_spanner_transaction.py b/tests/unit/adapters/test_spanner/test_spanner_transaction.py new file mode 100644 index 000000000..e27d4f64a --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_transaction.py @@ -0,0 +1,115 @@ +"""Unit tests for Spanner transaction retry closures.""" + +from typing import Any +from unittest.mock import MagicMock + +from google.api_core.exceptions import Aborted + +from sqlspec.adapters.spanner._typing import SpannerTransaction +from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.driver import SpannerSyncDriver +from sqlspec.exceptions import DeadlockError + + +def _create_mock_database(attempts_before_success: int = 1) -> tuple[MagicMock, MagicMock]: + """Create a mock database that retries on Aborted exceptions.""" + mock_db = MagicMock() + mock_txn = MagicMock(spec=SpannerTransaction) + attempts = [0] + + def mock_run_in_transaction(callback: Any, *args: Any, **kwargs: Any) -> Any: + while True: + attempts[0] += 1 + try: + return callback(mock_txn, *args, **kwargs) + except Aborted: + if attempts[0] > attempts_before_success: + raise + continue + + mock_db.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) + return mock_db, mock_txn + + +def test_config_run_in_transaction_retry_on_aborted() -> None: + """Verify that config.run_in_transaction retries when callback raises Aborted.""" + mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) + config = SpannerSyncConfig(connection_config={"project_id": "test", "instance_id": "inst", "database_id": "db"}) + config._database = mock_db + + call_count = [0] + + def unit_of_work(driver: SpannerSyncDriver) -> str: + call_count[0] += 1 + if call_count[0] == 1: + raise Aborted("Concurrency conflict") + return "success" + + result = config.run_in_transaction(unit_of_work) + + assert result == "success" + assert call_count[0] == 2 + mock_db.run_in_transaction.assert_called_once() + + +def test_config_run_in_transaction_retry_on_deadlock_error_with_aborted_cause() -> None: + """Verify that config.run_in_transaction unwraps DeadlockError caused by Aborted.""" + mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) + config = SpannerSyncConfig(connection_config={"project_id": "test", "instance_id": "inst", "database_id": "db"}) + config._database = mock_db + + call_count = [0] + + def unit_of_work(driver: SpannerSyncDriver) -> str: + call_count[0] += 1 + if call_count[0] == 1: + abort_exc = Aborted("Lock conflict") + deadlock_exc = DeadlockError("transaction aborted") + deadlock_exc.__cause__ = abort_exc + raise deadlock_exc + return "success-after-deadlock" + + result = config.run_in_transaction(unit_of_work) + + assert result == "success-after-deadlock" + assert call_count[0] == 2 + mock_db.run_in_transaction.assert_called_once() + + +def test_driver_run_in_transaction_delegates_to_database() -> None: + """Verify that driver.run_in_transaction delegates to database when not in transaction.""" + mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) + mock_session = MagicMock() + mock_session._database = mock_db + mock_snapshot = MagicMock() + mock_snapshot._session = mock_session + + driver = SpannerSyncDriver(connection=mock_snapshot) + + call_count = [0] + + def unit_of_work(txn_driver: SpannerSyncDriver) -> str: + call_count[0] += 1 + if call_count[0] == 1: + raise Aborted("Retryable abort") + return "driver-delegated" + + result = driver.run_in_transaction(unit_of_work) + + assert result == "driver-delegated" + assert call_count[0] == 2 + mock_db.run_in_transaction.assert_called_once() + + +def test_driver_run_in_transaction_in_existing_transaction() -> None: + """Verify that driver.run_in_transaction runs directly when connection is already a transaction.""" + mock_txn = MagicMock(spec=SpannerTransaction) + driver = SpannerSyncDriver(connection=mock_txn) + + def unit_of_work(txn_driver: SpannerSyncDriver) -> str: + assert txn_driver is driver + return "direct-execution" + + result = driver.run_in_transaction(unit_of_work) + + assert result == "direct-execution" diff --git a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py new file mode 100644 index 000000000..9ef0df154 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py @@ -0,0 +1,42 @@ +"""Unit tests for Decimal and INTERVAL parameter type inference.""" + +from datetime import timedelta +from decimal import Decimal + +from google.cloud.spanner_v1.types.type import TypeCode + +from sqlspec.adapters.spanner.type_converter import infer_spanner_param_types +from sqlspec.core import TypedParameter + + +def test_infer_decimal_param_types() -> None: + """Verify that non-null Decimal parameters infer as NUMERIC.""" + params = {"price": Decimal("19.99")} + types = infer_spanner_param_types(params) + assert "price" in types + assert types["price"].code == TypeCode.NUMERIC + + +def test_infer_decimal_array_param_types() -> None: + """Verify that sequences of Decimal parameters infer as Array(NUMERIC).""" + params = {"prices": [Decimal("19.99"), Decimal("29.99")]} + types = infer_spanner_param_types(params) + assert "prices" in types + assert types["prices"].code == TypeCode.ARRAY + assert types["prices"].array_element_type.code == TypeCode.NUMERIC + + +def test_infer_timedelta_param_types() -> None: + """Verify that timedelta parameters infer as INTERVAL.""" + params = {"duration": timedelta(days=1, hours=2)} + types = infer_spanner_param_types(params) + assert "duration" in types + assert types["duration"].code == TypeCode.INTERVAL + + +def test_null_timedelta_param_types() -> None: + """Verify that null TypedParameter with timedelta resolves to INTERVAL.""" + params = {"duration": TypedParameter(None, timedelta)} + types = infer_spanner_param_types(params) + assert "duration" in types + assert types["duration"].code == TypeCode.INTERVAL diff --git a/tests/unit/adapters/test_spanner/test_spanner_vector.py b/tests/unit/adapters/test_spanner/test_spanner_vector.py new file mode 100644 index 000000000..ab0390adb --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_vector.py @@ -0,0 +1,48 @@ +"""Unit tests for Spanner FLOAT32 and ARRAY vector parameter typing.""" + +from google.cloud.spanner_v1.types.type import TypeCode + +from sqlspec.adapters.spanner.type_converter import coerce_params_for_spanner, infer_spanner_param_types +from sqlspec.core import TypedParameter + + +def test_infer_vector_param_types_array_float32() -> None: + """Verify that TypedParameter with ARRAY infers as Array(FLOAT32).""" + params = {"embedding": TypedParameter([0.1, 0.2, 0.3], "ARRAY")} + types = infer_spanner_param_types(params) + assert "embedding" in types + assert types["embedding"].code == TypeCode.ARRAY + assert types["embedding"].array_element_type.code == TypeCode.FLOAT32 + + +def test_infer_vector_param_types_vector_alias() -> None: + """Verify that TypedParameter with VECTOR infers as Array(FLOAT32).""" + params = {"embedding": TypedParameter([0.1, 0.2, 0.3], "VECTOR")} + types = infer_spanner_param_types(params) + assert "embedding" in types + assert types["embedding"].code == TypeCode.ARRAY + assert types["embedding"].array_element_type.code == TypeCode.FLOAT32 + + +def test_infer_scalar_float32() -> None: + """Verify that TypedParameter with FLOAT32 infers as FLOAT32.""" + params = {"score": TypedParameter(1.25, "FLOAT32")} + types = infer_spanner_param_types(params) + assert "score" in types + assert types["score"].code == TypeCode.FLOAT32 + + +def test_coerce_vector_params() -> None: + """Verify that TypedParameter vector is unwrapped into a list of floats.""" + params = {"embedding": TypedParameter((0.1, 0.2, 0.3), "ARRAY")} + coerced = coerce_params_for_spanner(params) + assert coerced is not None + assert coerced["embedding"] == [0.1, 0.2, 0.3] + + +def test_null_float32_param_type() -> None: + """Verify that NULL TypedParameter with FLOAT32 resolves to param_types.FLOAT32.""" + params = {"score": TypedParameter(None, "FLOAT32")} + types = infer_spanner_param_types(params) + assert "score" in types + assert types["score"].code == TypeCode.FLOAT32 From 0b76ef1cb7c2d66da62dcf7c67321b8473cb3cab Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 20:33:49 +0000 Subject: [PATCH 02/15] fix(spanner): resolve quality typing errors in driver, config, and tests --- sqlspec/adapters/bigquery/core.py | 6 ++++-- sqlspec/adapters/spanner/config.py | 7 +++++-- sqlspec/adapters/spanner/core.py | 2 +- sqlspec/adapters/spanner/driver.py | 3 ++- tests/unit/adapters/test_spanner/test_spanner_json.py | 7 ++++--- tests/unit/adapters/test_spanner/test_spanner_stores.py | 7 ++++++- .../adapters/test_spanner/test_spanner_transaction.py | 8 ++++---- 7 files changed, 26 insertions(+), 14 deletions(-) diff --git a/sqlspec/adapters/bigquery/core.py b/sqlspec/adapters/bigquery/core.py index 1ea8205f1..c7d323eff 100644 --- a/sqlspec/adapters/bigquery/core.py +++ b/sqlspec/adapters/bigquery/core.py @@ -254,7 +254,7 @@ def create_parameters(parameters: Any, json_serializer: "Callable[[Any], str] | if _is_query_parameter(value): bq_parameters.append(cast("BigQueryParam", value)) continue - declared_type: type[Any] | None = None + declared_type: type[Any] | str | None = None if type(value) is TypedParameter: declared_type = value.original_type actual_value = value.value @@ -1007,7 +1007,9 @@ def _load_bigquery_module() -> Any: return _BIGQUERY_MODULE -def _query_parameter_type(value: Any, declared_type: "type[Any] | None" = None) -> "tuple[str | None, str | None]": +def _query_parameter_type( + value: Any, declared_type: "type[Any] | str | None" = None +) -> "tuple[str | None, str | None]": """Determine BigQuery parameter type from Python value. Args: diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index 8ca11ed35..a5bdabd57 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -414,6 +414,7 @@ def create_connection(self) -> SpannerConnection: return cast("SpannerConnection", self.get_database().snapshot(multi_use=True)) # type: ignore[no-untyped-call] def _create_pool(self) -> "AbstractSessionPool": + from sqlspec.adapters.spanner._typing import SpannerAbstractSessionPool as AbstractSessionPool from sqlspec.adapters.spanner._typing import SpannerBurstyPool as BurstyPool from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool @@ -425,6 +426,7 @@ def _create_pool(self) -> "AbstractSessionPool": raise ImproperConfigurationError(msg) raw_pool_type = self.connection_config.get("pool_type") + pool_type: type[AbstractSessionPool | PingingPool] if raw_pool_type is None or raw_pool_type == "multiplexed": pool_type = PingingPool else: @@ -637,7 +639,7 @@ def _callback(spanner_transaction: Any, *cb_args: Any, **cb_kwargs: Any) -> Any: raise cause from exc raise - return database.run_in_transaction(_callback, *args, **kwargs) + return cast("Any", database).run_in_transaction(_callback, *args, **kwargs) def execute_partitioned_dml( self, @@ -689,7 +691,8 @@ def execute_partitioned_dml( if effective_request_options is not None: call_kwargs["request_options"] = effective_request_options - return database.execute_partitioned_dml(sql, **call_kwargs) + result = cast("Any", database).execute_partitioned_dml(sql, **call_kwargs) + return int(result) def _session_driver_features( self, diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index eecada431..e12c9f631 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -320,7 +320,7 @@ def _unwrap_spanner_json_object(val: Any) -> Any: return [_unwrap_spanner_json_object(item) for item in array_val] if array_val is not None else [] if getattr(val, "_is_scalar_value", False): return getattr(val, "_simple_value", None) - return {k: _unwrap_spanner_json_object(v) for k, v in val.items()} + return {k: _unwrap_spanner_json_object(v) for k, v in cast("dict[str, Any]", val).items()} if isinstance(val, dict): return {k: _unwrap_spanner_json_object(v) for k, v in val.items()} if isinstance(val, (list, tuple)): diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index f78903979..2c9329e97 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -396,7 +396,8 @@ def execute_partitioned_dml( if effective_request_options is not None: call_kwargs["request_options"] = effective_request_options - return database.execute_partitioned_dml(sql, **call_kwargs) + result = database.execute_partitioned_dml(sql, **call_kwargs) + return int(result) def execute( self, diff --git a/tests/unit/adapters/test_spanner/test_spanner_json.py b/tests/unit/adapters/test_spanner/test_spanner_json.py index 824a80d3d..8f2dd71e1 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_json.py +++ b/tests/unit/adapters/test_spanner/test_spanner_json.py @@ -1,5 +1,6 @@ """Unit tests for Spanner JsonObject direct unwrapping optimization.""" +from typing import Any, cast from unittest.mock import MagicMock from google.cloud.spanner_v1.data_types import JsonObject @@ -11,13 +12,13 @@ class MonitoredJsonObject(JsonObject): """JsonObject subclass that tracks calls to serialize().""" - def __init__(self, *args, **kwargs) -> None: - super().__init__(*args, **kwargs) + def __init__(self, *args: Any, **kwargs: Any) -> None: + cast("Any", super()).__init__(*args, **kwargs) self.serialize_called = False def serialize(self) -> str | None: self.serialize_called = True - return super().serialize() + return cast("str | None", cast("Any", super()).serialize()) def test_convert_json_row_value_unwraps_without_serialize() -> None: diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py index fe29d8021..b799bd90d 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -45,7 +45,12 @@ def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: mock_driver = MagicMock(spec=SpannerSyncDriver) mock_result = MagicMock() mock_result.rowcount = 1 - mock_driver.execute.side_effect = lambda sql, *a, **kw: (executed_sqls.append(str(sql)), mock_result)[1] + + def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: + executed_sqls.append(str(sql)) + return mock_result + + mock_driver.execute.side_effect = mock_execute return func(mock_driver, *args, **kwargs) config.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) diff --git a/tests/unit/adapters/test_spanner/test_spanner_transaction.py b/tests/unit/adapters/test_spanner/test_spanner_transaction.py index e27d4f64a..313c390bb 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_transaction.py +++ b/tests/unit/adapters/test_spanner/test_spanner_transaction.py @@ -1,6 +1,6 @@ """Unit tests for Spanner transaction retry closures.""" -from typing import Any +from typing import Any, cast from unittest.mock import MagicMock from google.api_core.exceptions import Aborted @@ -42,7 +42,7 @@ def test_config_run_in_transaction_retry_on_aborted() -> None: def unit_of_work(driver: SpannerSyncDriver) -> str: call_count[0] += 1 if call_count[0] == 1: - raise Aborted("Concurrency conflict") + raise cast("Any", Aborted)("Concurrency conflict") return "success" result = config.run_in_transaction(unit_of_work) @@ -63,7 +63,7 @@ def test_config_run_in_transaction_retry_on_deadlock_error_with_aborted_cause() def unit_of_work(driver: SpannerSyncDriver) -> str: call_count[0] += 1 if call_count[0] == 1: - abort_exc = Aborted("Lock conflict") + abort_exc = cast("Any", Aborted)("Lock conflict") deadlock_exc = DeadlockError("transaction aborted") deadlock_exc.__cause__ = abort_exc raise deadlock_exc @@ -91,7 +91,7 @@ def test_driver_run_in_transaction_delegates_to_database() -> None: def unit_of_work(txn_driver: SpannerSyncDriver) -> str: call_count[0] += 1 if call_count[0] == 1: - raise Aborted("Retryable abort") + raise cast("Any", Aborted)("Retryable abort") return "driver-delegated" result = driver.run_in_transaction(unit_of_work) From 8b35b88f7f87ddf77401921f63b13fb036e91881 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 21:05:13 +0000 Subject: [PATCH 03/15] fix(spanner): resolve adk delete params, litestar bytes decoding, and consumed feature keys --- sqlspec/adapters/spanner/adk/store.py | 11 ++++++++--- sqlspec/adapters/spanner/litestar/store.py | 2 +- .../adapters/_shared/_driver_type_system.py | 1 + 3 files changed, 10 insertions(+), 4 deletions(-) diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index 9c347d8fa..d3a370b2c 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -443,13 +443,18 @@ def _delete_session(self, app_name: str, user_id: str, session_id: str) -> None: ) delete_events_sql = f"DELETE FROM {self._events_table} WHERE session_id = @session_id{shard_clause}" delete_session_sql = f"DELETE FROM {self._session_table} WHERE app_name = @app_name AND user_id = @user_id AND id = @session_id{shard_clause}" - params = {"app_name": app_name, "user_id": user_id, "session_id": session_id} - types = { + delete_events_params = {"session_id": session_id} + delete_events_types = {"session_id": SPANNER_PARAM_TYPES.STRING} + delete_session_params = {"app_name": app_name, "user_id": user_id, "session_id": session_id} + delete_session_types = { "app_name": SPANNER_PARAM_TYPES.STRING, "user_id": SPANNER_PARAM_TYPES.STRING, "session_id": SPANNER_PARAM_TYPES.STRING, } - self._run_write([(delete_events_sql, params, types), (delete_session_sql, params, types)]) + self._run_write([ + (delete_events_sql, delete_events_params, delete_events_types), + (delete_session_sql, delete_session_params, delete_session_types), + ]) def _append_event_and_update_state( self, diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index 31db32c85..18c0d7d16 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -151,7 +151,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if result is None: return None - data = result.get("data") + data = spanner_to_bytes(result.get("data")) expires_at = self._timestamp_to_datetime(result.get("expires_at")) if renew_for is not None and expires_at is not None: diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 540f43c86..aa561ab7c 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -248,6 +248,7 @@ class SourceEquivalenceCase: "enable_events", "events_backend", "enable_batch_write_api", + "query_options", ), "sqlite": ( "enable_custom_adapters", From 3b5e46f25da0a34a319e86e019efd921e50593d3 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 21:42:10 +0000 Subject: [PATCH 04/15] fix(spanner): bind native JSON payloads in ADK store and align rows_affected in litestar store --- sqlspec/adapters/spanner/adk/store.py | 61 ++++++++++++++++------ sqlspec/adapters/spanner/litestar/store.py | 18 ++++--- 2 files changed, 55 insertions(+), 24 deletions(-) diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index d3a370b2c..f620b42ac 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -9,6 +9,7 @@ from sqlspec.adapters.spanner._typing import SpannerNotFound as NotFound from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.core import _unwrap_spanner_json_object from sqlspec.config import ADKConfig from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore @@ -268,22 +269,30 @@ def _metadata_param_types(self) -> "dict[str, Any]": return {"key": SPANNER_PARAM_TYPES.STRING, "value": SPANNER_PARAM_TYPES.STRING} def _decode_state(self, raw: Any) -> Any: + if raw is None: + return None if isinstance(raw, str): - return from_json(raw) - return raw + try: + return from_json(raw) + except Exception: + return raw + return _unwrap_spanner_json_object(raw) def _decode_json(self, raw: Any) -> Any: if raw is None: return None if isinstance(raw, str): - return from_json(raw) - return raw + try: + return from_json(raw) + except Exception: + return raw + return _unwrap_spanner_json_object(raw) def _create_session( self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None ) -> StoredSession: - state_json = to_json(state) - params: dict[str, Any] = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_json} + state_payload = _to_spanner_json_payload(state) + params: dict[str, Any] = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_payload} columns = "id, app_name, user_id, state, create_time, update_time" values = "@id, @app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()" if self._owner_id_column_name: @@ -365,7 +374,7 @@ def _get_session( return record def _update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: - params = {"app_name": app_name, "user_id": user_id, "id": session_id, "state": to_json(state)} + params = {"app_name": app_name, "user_id": user_id, "id": session_id, "state": _to_spanner_json_payload(state)} json_type = _json_param_type() sql = f""" UPDATE {self._session_table} @@ -491,7 +500,7 @@ def _append_event_and_update_state( "session_id": event_record["session_id"], "invocation_id": event_record["invocation_id"], "timestamp": event_record["timestamp"], - "event_data": to_json(event_record["event_data"]), + "event_data": _to_spanner_json_payload(event_record["event_data"]), } insert_sql = f""" INSERT INTO {self._events_table} (id, app_name, user_id, session_id, invocation_id, timestamp, event_data) @@ -503,7 +512,7 @@ def _append_event_and_update_state( "app_name": app_name, "user_id": user_id, "id": session_id, - "state": to_json(state), + "state": _to_spanner_json_payload(state), } update_sql = f""" UPDATE {self._session_table} @@ -532,7 +541,7 @@ def _append_event_and_update_state( INSERT OR UPDATE {self._app_state_table} (app_name, state, update_time) VALUES (@app_name, @state, PENDING_COMMIT_TIMESTAMP()) """, - {"app_name": app_name, "state": to_json(app_state)}, + {"app_name": app_name, "state": _to_spanner_json_payload(app_state)}, self._app_state_param_types(), )) if user_state is not None: @@ -541,7 +550,7 @@ def _append_event_and_update_state( INSERT OR UPDATE {self._user_state_table} (app_name, user_id, state, update_time) VALUES (@app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP()) """, - {"app_name": app_name, "user_id": user_id, "state": to_json(user_state)}, + {"app_name": app_name, "user_id": user_id, "state": _to_spanner_json_payload(user_state)}, self._user_state_param_types(), )) @@ -561,7 +570,7 @@ def _insert_event(self, event_record: "StoredEvent") -> None: "session_id": event_record["session_id"], "invocation_id": event_record["invocation_id"], "timestamp": event_record["timestamp"], - "event_data": to_json(event_record["event_data"]), + "event_data": _to_spanner_json_payload(event_record["event_data"]), } insert_sql = f""" INSERT INTO {self._events_table} (id, app_name, user_id, session_id, invocation_id, timestamp, event_data) @@ -628,7 +637,7 @@ def _delete_expired_events(self, before: datetime, app_name: "str | None" = None sql += " AND app_name = @app_name" params["app_name"] = app_name result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rowcount", 0)) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._session_table} WHERE update_time < @updated_before" @@ -637,7 +646,7 @@ def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" sql += " AND app_name = @app_name" params["app_name"] = app_name result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rowcount", 0)) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._user_state_table} WHERE update_time < @updated_before" @@ -646,7 +655,7 @@ def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | No sql += " AND app_name = @app_name" params["app_name"] = app_name result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rowcount", 0)) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = @app_name LIMIT 1" @@ -676,7 +685,9 @@ def _upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: INSERT OR UPDATE {self._app_state_table} (app_name, state, update_time) VALUES (@app_name, @state, PENDING_COMMIT_TIMESTAMP()) """ - self._run_write([(sql, {"app_name": app_name, "state": to_json(state)}, self._app_state_param_types())]) + self._run_write([ + (sql, {"app_name": app_name, "state": _to_spanner_json_payload(state)}, self._app_state_param_types()) + ]) def _upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: sql = f""" @@ -684,7 +695,11 @@ def _upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any] VALUES (@app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP()) """ self._run_write([ - (sql, {"app_name": app_name, "user_id": user_id, "state": to_json(state)}, self._user_state_param_types()) + ( + sql, + {"app_name": app_name, "user_id": user_id, "state": _to_spanner_json_payload(state)}, + self._user_state_param_types(), + ) ]) def _get_metadata(self, key: str) -> "str | None": @@ -1151,6 +1166,18 @@ def _rows_to_records(self, rows: "list[Any]") -> "list[StoredMemory]": ] +def _to_spanner_json_payload(value: Any) -> Any: + """Prepare a value for Spanner JSON column parameter binding.""" + if value is None: + return None + if isinstance(value, (str, bytes)): + try: + return from_json(value) + except Exception: + return value + return value + + def _json_param_type() -> Any: try: return SPANNER_PARAM_TYPES.JSON diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index 18c0d7d16..664dcb4f3 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -111,11 +111,10 @@ def _timestamp_to_datetime(self, ts: "datetime | None") -> "datetime | None": def _build_params( self, key: str, expires_at: "datetime | None" = None, data: "bytes | None" = None ) -> "dict[str, Any]": - return { - "session_id": key, - "data": bytes_to_spanner(data), - "expires_at": self._datetime_to_timestamp(expires_at), - } + params: dict[str, Any] = {"session_id": key, "expires_at": self._datetime_to_timestamp(expires_at)} + if data is not None: + params["data"] = bytes_to_spanner(data) + return params def _get_param_types( self, session_id: bool = True, expires_at: bool = False, data: bool = False @@ -189,7 +188,9 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No def _job(driver: "SpannerSyncDriver") -> None: result = driver.execute(update_sql, params) - if not getattr(result, "rowcount", None): + rows_affected = getattr(result, "rows_affected", None) + has_rows = rows_affected > 0 if isinstance(rows_affected, int) else bool(getattr(result, "rowcount", None)) + if not has_rows: driver.execute(insert_sql, params) self._config.run_in_transaction(_job) @@ -239,7 +240,10 @@ def _delete_expired(self) -> int: WHERE expires_at IS NOT NULL AND expires_at <= CURRENT_TIMESTAMP() """ result = self._config.run_in_transaction(lambda driver: driver.execute(sql)) - return cast("int", getattr(result, "rowcount", 0)) + rows_affected = getattr(result, "rows_affected", None) + if isinstance(rows_affected, int): + return rows_affected + return int(getattr(result, "rowcount", 0)) def _create_table(self) -> None: database = self._config.get_database() From 2a59d1941a9b8c84c942c8e9efddaa39e22fd4e8 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Fri, 25 Sep 2026 01:01:07 +0000 Subject: [PATCH 05/15] refactor(spanner): route store operations through provide_session and remove ad-hoc methods --- sqlspec/adapters/spanner/adk/store.py | 29 +++-- sqlspec/adapters/spanner/config.py | 92 +------------- sqlspec/adapters/spanner/driver.py | 46 ------- sqlspec/adapters/spanner/litestar/store.py | 25 ++-- .../test_spanner/test_litestar_store.py | 50 +++++--- .../test_spanner_partitioned_dml.py | 48 +------- .../test_spanner/test_spanner_stores.py | 66 +++++----- .../test_spanner/test_spanner_transaction.py | 115 ------------------ 8 files changed, 97 insertions(+), 374 deletions(-) delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_transaction.py diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index f620b42ac..2673c9493 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -20,7 +20,6 @@ from collections.abc import Sequence from sqlspec.adapters.spanner._typing import SpannerDatabase as Database - from sqlspec.adapters.spanner.driver import SpannerSyncDriver from sqlspec.extensions.adk import SessionOrderBy, StoredMemory __all__ = ("SpannerADKConfig", "SpannerADKRetentionConfig", "SpannerSyncADKMemoryStore", "SpannerSyncADKStore") @@ -225,12 +224,10 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - def _job(driver: "SpannerSyncDriver") -> None: + with self._config.provide_session() as driver: for sql, params, _ in statements: driver.execute(sql, params) - self._config.run_in_transaction(_job) - def _session_param_types(self, include_owner: bool) -> "dict[str, Any]": json_type = _json_param_type() types: dict[str, Any] = { @@ -636,8 +633,9 @@ def _delete_expired_events(self, before: datetime, app_name: "str | None" = None if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + with self._config.provide_session() as driver: + result = driver.execute(sql, params) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._session_table} WHERE update_time < @updated_before" @@ -645,8 +643,9 @@ def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + with self._config.provide_session() as driver: + result = driver.execute(sql, params) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._user_state_table} WHERE update_time < @updated_before" @@ -654,8 +653,9 @@ def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | No if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + with self._config.provide_session() as driver: + result = driver.execute(sql, params) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = @app_name LIMIT 1" @@ -899,15 +899,14 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - def _job(driver: "SpannerSyncDriver") -> None: + with self._config.provide_session() as driver: for sql, params, _ in statements: driver.execute(sql, params) - self._config.run_in_transaction(_job) - def _execute_update(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> int: - result = self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) - return int(getattr(result, "rowcount", 0)) + with self._config.provide_session() as driver: + result = driver.execute(sql, params) + return int(getattr(result, "rowcount", 0)) def _memory_param_types(self, include_owner: bool) -> "dict[str, Any]": types: dict[str, Any] = { diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index a5bdabd57..a678205e9 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -9,9 +9,8 @@ from sqlspec.adapters.spanner._typing import SpannerTransactionType as TransactionType from sqlspec.adapters.spanner.core import apply_driver_features, default_statement_config from sqlspec.adapters.spanner.driver import SpannerSessionContext, SpannerSyncDriver -from sqlspec.adapters.spanner.type_converter import coerce_params_for_spanner, infer_spanner_param_types from sqlspec.config import SyncDatabaseConfig -from sqlspec.core import SQL, TypeCoercionCapabilities +from sqlspec.core import TypeCoercionCapabilities from sqlspec.driver import SyncPoolConnectionContext, SyncPoolSessionFactory from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints @@ -605,95 +604,6 @@ def provide_read_session( **kwargs, ) - def run_in_transaction(self, func: "Callable[[SpannerSyncDriver], Any]", *args: Any, **kwargs: Any) -> Any: - """Execute a unit of work inside a transaction, retrying on abort. - - Args: - func: Callback taking a prepared SpannerSyncDriver and returning a result. - *args: Positional arguments passed to the callback. - **kwargs: Keyword arguments passed to the callback. - - Returns: - The return value of the callback. - """ - database = self.get_database() - - def _callback(spanner_transaction: Any, *cb_args: Any, **cb_kwargs: Any) -> Any: - driver = SpannerSyncDriver( - connection=spanner_transaction, - statement_config=self.statement_config, - driver_features=self.driver_features, - ) - prepared_driver = self._prepare_driver(driver) - call_args = cb_args or args - call_kwargs = cb_kwargs or kwargs - try: - return func(prepared_driver, *call_args, **call_kwargs) - except Exception as exc: - cause = getattr(exc, "__cause__", None) - from google.api_core import exceptions as api_exceptions - - if isinstance(exc, api_exceptions.Aborted): - raise - if cause is not None and isinstance(cause, api_exceptions.Aborted): - raise cause from exc - raise - - return cast("Any", database).run_in_transaction(_callback, *args, **kwargs) - - def execute_partitioned_dml( - self, - statement: "SQL | str", - *parameters: Any, - query_options: Any = None, - request_options: Any = None, - exclude_txn_from_change_streams: bool = False, - **kwargs: Any, - ) -> int: - """Execute a Partitioned DML statement across database partitions. - - Args: - statement: The SQL string or SQL object to execute. - *parameters: Positional parameters or parameter mapping. - query_options: Optional Spanner QueryOptions. - request_options: Optional Spanner RequestOptions. - exclude_txn_from_change_streams: Whether to exclude the transaction from change streams. - **kwargs: Additional keyword arguments or parameters. - - Returns: - The number of affected rows. - """ - database = self.get_database() - if isinstance(statement, SQL): - sql_statement = statement - else: - sql_statement = SQL(statement, *parameters, statement_config=self.statement_config, **kwargs) - - sql, raw_params = sql_statement.compile() - params = raw_params if isinstance(raw_params, dict) else None - coerced_params = coerce_params_for_spanner( - params, - json_serializer=self.driver_features.get("json_serializer"), - enable_uuid_conversion=self.driver_features.get("enable_uuid_conversion", True), - ) - param_types = infer_spanner_param_types(params) - - effective_request_options = request_options or self.driver_features.get("request_options") - effective_query_options = query_options or self.driver_features.get("query_options") - - call_kwargs: dict[str, Any] = { - "params": coerced_params, - "param_types": param_types, - "exclude_txn_from_change_streams": exclude_txn_from_change_streams, - } - if effective_query_options is not None: - call_kwargs["query_options"] = effective_query_options - if effective_request_options is not None: - call_kwargs["request_options"] = effective_request_options - - result = cast("Any", database).execute_partitioned_dml(sql, **call_kwargs) - return int(result) - def _session_driver_features( self, *, diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 2c9329e97..ba18c4556 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -298,52 +298,6 @@ def _get_database(self) -> Any: return database return None - def run_in_transaction(self, func: "Callable[[SpannerSyncDriver], Any]", *args: Any, **kwargs: Any) -> Any: - """Execute a unit of work inside a transaction, retrying on abort. - - If the driver is already bound to a SpannerTransaction, the callable - is executed directly. Otherwise, work is delegated to the database - transaction retry runner. - - Args: - func: Callback taking this driver or a transaction-bound driver. - *args: Positional arguments passed to the callback. - **kwargs: Keyword arguments passed to the callback. - - Returns: - The return value of the callback. - """ - if isinstance(self.connection, SpannerTransaction): - return func(self, *args, **kwargs) - - database = self._get_database() - if database is None: - msg = "run_in_transaction requires an active database or SpannerTransaction context." - raise SQLConversionError(msg) - - def _callback(spanner_transaction: Any, *cb_args: Any, **cb_kwargs: Any) -> Any: - driver = SpannerSyncDriver( - connection=spanner_transaction, - statement_config=self.statement_config, - driver_features=self.driver_features, - ) - driver._config = self._config - call_args = cb_args or args - call_kwargs = cb_kwargs or kwargs - try: - return func(driver, *call_args, **call_kwargs) - except Exception as exc: - cause = getattr(exc, "__cause__", None) - from google.api_core import exceptions as api_exceptions - - if isinstance(exc, api_exceptions.Aborted): - raise - if cause is not None and isinstance(cause, api_exceptions.Aborted): - raise cause from exc - raise - - return database.run_in_transaction(_callback, *args, **kwargs) - def execute_partitioned_dml( self, statement: "SQL | str", diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index 664dcb4f3..1996a971a 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -17,7 +17,6 @@ from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction from sqlspec.adapters.spanner.config import SpannerSyncConfig - from sqlspec.adapters.spanner.driver import SpannerSyncDriver class _DatabaseProtocol(Protocol): def run_in_transaction(self, func: "Callable[[Transaction], Any]") -> Any: ... @@ -163,7 +162,8 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if self._shard_count > 1: update_sql = f"{update_sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = self._build_params(key, new_expires) - self._config.run_in_transaction(lambda driver: driver.execute(update_sql, params)) + with self._config.provide_session() as driver: + driver.execute(update_sql, params) return spanner_to_bytes(data) @@ -186,25 +186,25 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No VALUES (@session_id, @data, @expires_at, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()) """ - def _job(driver: "SpannerSyncDriver") -> None: + with self._config.provide_session() as driver: result = driver.execute(update_sql, params) rows_affected = getattr(result, "rows_affected", None) has_rows = rows_affected > 0 if isinstance(rows_affected, int) else bool(getattr(result, "rowcount", None)) if not has_rows: driver.execute(insert_sql, params) - self._config.run_in_transaction(_job) - def _delete(self, key: str) -> None: sql = f"DELETE FROM {self._table_name} WHERE session_id = @session_id" if self._shard_count > 1: sql = f"{sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = {"session_id": key} - self._config.run_in_transaction(lambda driver: driver.execute(sql, params)) + with self._config.provide_session() as driver: + driver.execute(sql, params) def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name} WHERE TRUE" - self._config.run_in_transaction(lambda driver: driver.execute(sql)) + with self._config.provide_session() as driver: + driver.execute(sql) def _exists(self, key: str) -> bool: sql = f""" @@ -239,11 +239,12 @@ def _delete_expired(self) -> int: DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= CURRENT_TIMESTAMP() """ - result = self._config.run_in_transaction(lambda driver: driver.execute(sql)) - rows_affected = getattr(result, "rows_affected", None) - if isinstance(rows_affected, int): - return rows_affected - return int(getattr(result, "rowcount", 0)) + with self._config.provide_session() as driver: + result = driver.execute(sql) + rows_affected = getattr(result, "rows_affected", None) + if isinstance(rows_affected, int): + return rows_affected + return int(getattr(result, "rowcount", 0)) def _create_table(self) -> None: database = self._config.get_database() diff --git a/tests/unit/adapters/test_spanner/test_litestar_store.py b/tests/unit/adapters/test_spanner/test_litestar_store.py index 8f4d7525e..4aff468fb 100644 --- a/tests/unit/adapters/test_spanner/test_litestar_store.py +++ b/tests/unit/adapters/test_spanner/test_litestar_store.py @@ -4,55 +4,67 @@ from sqlspec.adapters.spanner.litestar import SpannerSyncStore -def _mock_database() -> MagicMock: - """Create a mock database that captures run_in_transaction calls.""" - db = MagicMock() - db.run_in_transaction = MagicMock(side_effect=lambda func: func(MagicMock())) - return db - +def test_set_uses_session() -> None: + """Verify _set uses config.provide_session for write operations.""" + driver = MagicMock() + driver.execute.return_value = MagicMock(rows_affected=1) + cm = _context_manager_yielding(driver) -def test_set_uses_run_in_transaction() -> None: - """Verify _set uses config.run_in_transaction for write operations.""" config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} + config.provide_session.return_value = cm store = SpannerSyncStore(config) store._set("s1", b"data", None) # pyright: ignore - config.run_in_transaction.assert_called_once() + config.provide_session.assert_called_once() -def test_delete_uses_run_in_transaction() -> None: - """Verify _delete uses config.run_in_transaction for write operations.""" +def test_delete_uses_session() -> None: + """Verify _delete uses config.provide_session for write operations.""" + driver = MagicMock() + cm = _context_manager_yielding(driver) + config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} + config.provide_session.return_value = cm store = SpannerSyncStore(config) store._delete("s1") # pyright: ignore - config.run_in_transaction.assert_called_once() + config.provide_session.assert_called_once() -def test_delete_all_uses_run_in_transaction() -> None: - """Verify _delete_all uses config.run_in_transaction for write operations.""" +def test_delete_all_uses_session() -> None: + """Verify _delete_all uses config.provide_session for write operations.""" + driver = MagicMock() + cm = _context_manager_yielding(driver) + config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} + config.provide_session.return_value = cm store = SpannerSyncStore(config) store._delete_all() # pyright: ignore - config.run_in_transaction.assert_called_once() + config.provide_session.assert_called_once() -def test_delete_expired_uses_run_in_transaction() -> None: - """Verify _delete_expired uses config.run_in_transaction for write operations.""" +def test_delete_expired_uses_session() -> None: + """Verify _delete_expired uses config.provide_session for write operations.""" + driver = MagicMock() + driver.execute.return_value = MagicMock(rows_affected=3) + cm = _context_manager_yielding(driver) + config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} + config.provide_session.return_value = cm store = SpannerSyncStore(config) - store._delete_expired() # pyright: ignore + result = store._delete_expired() # pyright: ignore - config.run_in_transaction.assert_called_once() + config.provide_session.assert_called_once() + assert result == 3 def _context_manager_yielding(value: Any) -> Any: diff --git a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py index cd42e46c9..4ebb55cd2 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py +++ b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py @@ -1,10 +1,9 @@ -"""Unit tests for Spanner Partitioned DML execution on driver and config.""" +"""Unit tests for Spanner Partitioned DML execution on driver.""" -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock from google.cloud.spanner_v1.types.type import TypeCode -from sqlspec.adapters.spanner.config import SpannerSyncConfig from sqlspec.adapters.spanner.core import default_statement_config from sqlspec.adapters.spanner.driver import SpannerSyncDriver @@ -53,21 +52,6 @@ def test_driver_execute_partitioned_dml_with_parameters() -> None: assert kwargs["param_types"]["limit"].code == TypeCode.INT64 -def test_config_execute_partitioned_dml() -> None: - """Verify that config.execute_partitioned_dml delegates to get_database().execute_partitioned_dml.""" - config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 100 - - with patch.object(config, "get_database", return_value=mock_db): - rows = config.execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") - - assert rows == 100 - mock_db.execute_partitioned_dml.assert_called_once() - sql = mock_db.execute_partitioned_dml.call_args[0][0] - assert "DELETE FROM large_table WHERE TRUE" in sql - - def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: """Verify executing partitioned DML with SQL object, query options, and request options.""" from sqlspec.core import SQL @@ -116,31 +100,3 @@ def test_driver_execute_partitioned_dml_no_database_raises() -> None: with pytest.raises(SQLConversionError, match="Could not resolve Spanner database"): driver.execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") - - -def test_config_execute_partitioned_dml_with_parameters_and_options() -> None: - """Verify config partitioned DML forwards parameters, types, and options.""" - config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 25 - - mock_query_options = MagicMock() - mock_request_options = MagicMock() - - with patch.object(config, "get_database", return_value=mock_db): - rows = config.execute_partitioned_dml( - "UPDATE items SET status = :status WHERE id = :id", - {"status": "deleted", "id": 5}, - query_options=mock_query_options, - request_options=mock_request_options, - exclude_txn_from_change_streams=True, - ) - - assert rows == 25 - mock_db.execute_partitioned_dml.assert_called_once() - _, kwargs = mock_db.execute_partitioned_dml.call_args - assert kwargs["params"] == {"status": "deleted", "id": 5} - assert "status" in kwargs["param_types"] - assert kwargs["query_options"] is mock_query_options - assert kwargs["request_options"] is mock_request_options - assert kwargs["exclude_txn_from_change_streams"] is True diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py index b799bd90d..112ffdb3d 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -1,4 +1,4 @@ -"""Unit tests for ADK and Litestar store transaction routing.""" +"""Unit tests for ADK and Litestar store session routing.""" from typing import Any from unittest.mock import MagicMock @@ -9,20 +9,28 @@ from sqlspec.adapters.spanner.litestar import SpannerSyncStore -def test_adk_store_run_write_routes_through_config_run_in_transaction() -> None: - """Verify that SpannerSyncADKStore._run_write executes via config.run_in_transaction.""" +def _context_manager_yielding(value: Any) -> Any: + class _Ctx: + def __enter__(self) -> Any: + return value + + def __exit__(self, *_: Any) -> None: + pass + + return _Ctx() + + +def test_adk_store_run_write_routes_through_provide_session() -> None: + """Verify that SpannerSyncADKStore._run_write executes via config.provide_session.""" config = MagicMock(spec=SpannerSyncConfig) executed_statements: list[tuple[str, Any]] = [] - def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_driver.execute.side_effect = lambda sql, params=None, param_types=None, **kw: executed_statements.append(( - sql, - params, - )) - return func(mock_driver, *args, **kwargs) - - config.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_driver.execute.side_effect = lambda sql, params=None, param_types=None, **kw: executed_statements.append(( + sql, + params, + )) + config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) store = SpannerSyncADKStore(config=config) statements = [ @@ -31,41 +39,39 @@ def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: ] store._run_write(statements) - config.run_in_transaction.assert_called_once() + config.provide_session.assert_called_once() assert len(executed_statements) == 2 -def test_litestar_store_writes_route_through_config_run_in_transaction() -> None: - """Verify that SpannerSyncStore write operations execute via config.run_in_transaction.""" +def test_litestar_store_writes_route_through_provide_session() -> None: + """Verify that SpannerSyncStore write operations execute via config.provide_session.""" config = MagicMock(spec=SpannerSyncConfig) config.extension_config = {"litestar": {"session_table": "sessions"}} executed_sqls: list[str] = [] - def mock_run_in_transaction(func: Any, *args: Any, **kwargs: Any) -> Any: - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_result = MagicMock() - mock_result.rowcount = 1 - - def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: - executed_sqls.append(str(sql)) - return mock_result + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_result = MagicMock() + mock_result.rowcount = 1 + mock_result.rows_affected = 1 - mock_driver.execute.side_effect = mock_execute - return func(mock_driver, *args, **kwargs) + def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: + executed_sqls.append(str(sql)) + return mock_result - config.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) + mock_driver.execute.side_effect = mock_execute + config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) store = SpannerSyncStore(config=config) store._set("session_1", b"payload", expires_in=3600) - assert config.run_in_transaction.call_count == 1 + assert config.provide_session.call_count == 1 store._delete("session_1") - assert config.run_in_transaction.call_count == 2 + assert config.provide_session.call_count == 2 store._delete_all() - assert config.run_in_transaction.call_count == 3 + assert config.provide_session.call_count == 3 expired_count = store._delete_expired() - assert config.run_in_transaction.call_count == 4 + assert config.provide_session.call_count == 4 assert expired_count == 1 diff --git a/tests/unit/adapters/test_spanner/test_spanner_transaction.py b/tests/unit/adapters/test_spanner/test_spanner_transaction.py deleted file mode 100644 index 313c390bb..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_transaction.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Unit tests for Spanner transaction retry closures.""" - -from typing import Any, cast -from unittest.mock import MagicMock - -from google.api_core.exceptions import Aborted - -from sqlspec.adapters.spanner._typing import SpannerTransaction -from sqlspec.adapters.spanner.config import SpannerSyncConfig -from sqlspec.adapters.spanner.driver import SpannerSyncDriver -from sqlspec.exceptions import DeadlockError - - -def _create_mock_database(attempts_before_success: int = 1) -> tuple[MagicMock, MagicMock]: - """Create a mock database that retries on Aborted exceptions.""" - mock_db = MagicMock() - mock_txn = MagicMock(spec=SpannerTransaction) - attempts = [0] - - def mock_run_in_transaction(callback: Any, *args: Any, **kwargs: Any) -> Any: - while True: - attempts[0] += 1 - try: - return callback(mock_txn, *args, **kwargs) - except Aborted: - if attempts[0] > attempts_before_success: - raise - continue - - mock_db.run_in_transaction = MagicMock(side_effect=mock_run_in_transaction) - return mock_db, mock_txn - - -def test_config_run_in_transaction_retry_on_aborted() -> None: - """Verify that config.run_in_transaction retries when callback raises Aborted.""" - mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) - config = SpannerSyncConfig(connection_config={"project_id": "test", "instance_id": "inst", "database_id": "db"}) - config._database = mock_db - - call_count = [0] - - def unit_of_work(driver: SpannerSyncDriver) -> str: - call_count[0] += 1 - if call_count[0] == 1: - raise cast("Any", Aborted)("Concurrency conflict") - return "success" - - result = config.run_in_transaction(unit_of_work) - - assert result == "success" - assert call_count[0] == 2 - mock_db.run_in_transaction.assert_called_once() - - -def test_config_run_in_transaction_retry_on_deadlock_error_with_aborted_cause() -> None: - """Verify that config.run_in_transaction unwraps DeadlockError caused by Aborted.""" - mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) - config = SpannerSyncConfig(connection_config={"project_id": "test", "instance_id": "inst", "database_id": "db"}) - config._database = mock_db - - call_count = [0] - - def unit_of_work(driver: SpannerSyncDriver) -> str: - call_count[0] += 1 - if call_count[0] == 1: - abort_exc = cast("Any", Aborted)("Lock conflict") - deadlock_exc = DeadlockError("transaction aborted") - deadlock_exc.__cause__ = abort_exc - raise deadlock_exc - return "success-after-deadlock" - - result = config.run_in_transaction(unit_of_work) - - assert result == "success-after-deadlock" - assert call_count[0] == 2 - mock_db.run_in_transaction.assert_called_once() - - -def test_driver_run_in_transaction_delegates_to_database() -> None: - """Verify that driver.run_in_transaction delegates to database when not in transaction.""" - mock_db, _mock_txn = _create_mock_database(attempts_before_success=1) - mock_session = MagicMock() - mock_session._database = mock_db - mock_snapshot = MagicMock() - mock_snapshot._session = mock_session - - driver = SpannerSyncDriver(connection=mock_snapshot) - - call_count = [0] - - def unit_of_work(txn_driver: SpannerSyncDriver) -> str: - call_count[0] += 1 - if call_count[0] == 1: - raise cast("Any", Aborted)("Retryable abort") - return "driver-delegated" - - result = driver.run_in_transaction(unit_of_work) - - assert result == "driver-delegated" - assert call_count[0] == 2 - mock_db.run_in_transaction.assert_called_once() - - -def test_driver_run_in_transaction_in_existing_transaction() -> None: - """Verify that driver.run_in_transaction runs directly when connection is already a transaction.""" - mock_txn = MagicMock(spec=SpannerTransaction) - driver = SpannerSyncDriver(connection=mock_txn) - - def unit_of_work(txn_driver: SpannerSyncDriver) -> str: - assert txn_driver is driver - return "direct-execution" - - result = driver.run_in_transaction(unit_of_work) - - assert result == "direct-execution" From 3b26e0985fb7b39396212ea61506e3760a989e1b Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 01:38:07 +0000 Subject: [PATCH 06/15] fix(spanner): fix store base64 double-encoding, ADK memory JSON/null binding, and PDML error mapping --- sqlspec/adapters/spanner/adk/store.py | 55 ++++++++++-- sqlspec/adapters/spanner/driver.py | 7 +- sqlspec/adapters/spanner/litestar/store.py | 40 ++------- sqlspec/adapters/spanner/type_converter.py | 5 ++ .../test_spanner/test_spanner_stores.py | 85 +++++++++++++++++++ .../test_spanner_type_inference.py | 8 ++ 6 files changed, 160 insertions(+), 40 deletions(-) diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index 2673c9493..c5f7ef00b 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -11,6 +11,7 @@ from sqlspec.adapters.spanner.config import SpannerSyncConfig from sqlspec.adapters.spanner.core import _unwrap_spanner_json_object from sqlspec.config import ADKConfig +from sqlspec.core import TypedParameter from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore from sqlspec.protocols import SpannerParamTypesProtocol @@ -225,8 +226,8 @@ def _run_read( def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: with self._config.provide_session() as driver: - for sql, params, _ in statements: - driver.execute(sql, params) + for sql, params, types in statements: + driver.execute(sql, _prepare_spanner_write_params(params, types)) def _session_param_types(self, include_owner: bool) -> "dict[str, Any]": json_type = _json_param_type() @@ -900,13 +901,13 @@ def _run_read( def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: with self._config.provide_session() as driver: - for sql, params, _ in statements: - driver.execute(sql, params) + for sql, params, types in statements: + driver.execute(sql, _prepare_spanner_write_params(params, types)) def _execute_update(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> int: with self._config.provide_session() as driver: - result = driver.execute(sql, params) - return int(getattr(result, "rowcount", 0)) + result = driver.execute(sql, _prepare_spanner_write_params(params, types)) + return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) def _memory_param_types(self, include_owner: bool) -> "dict[str, Any]": types: dict[str, Any] = { @@ -931,8 +932,11 @@ def _decode_json(self, raw: Any) -> Any: if raw is None: return None if isinstance(raw, str): - return from_json(raw) - return raw + try: + return from_json(raw) + except Exception: + return raw + return _unwrap_spanner_json_object(raw) def _create_tables(self) -> None: if not self._enabled: @@ -1165,6 +1169,41 @@ def _rows_to_records(self, rows: "list[Any]") -> "list[StoredMemory]": ] +def _prepare_spanner_write_params(params: "dict[str, Any]", types: "dict[str, Any] | None") -> "dict[str, Any]": + """Prepare ADK write parameters for Spanner driver execution.""" + if not types: + return params + json_type = _json_param_type() + changed = False + prepared: dict[str, Any] = {} + for key, value in params.items(): + param_type = types.get(key) + if param_type == json_type: + if value is None: + prepared[key] = TypedParameter(None, dict) + changed = True + elif isinstance(value, (str, bytes)): + prepared[key] = _to_spanner_json_payload(value) + changed = True + else: + prepared[key] = value + elif value is None and param_type is not None: + if param_type == SPANNER_PARAM_TYPES.STRING: + prepared[key] = TypedParameter(None, str) + changed = True + elif param_type == SPANNER_PARAM_TYPES.TIMESTAMP: + prepared[key] = TypedParameter(None, datetime) + changed = True + elif param_type == SPANNER_PARAM_TYPES.INT64: + prepared[key] = TypedParameter(None, int) + changed = True + else: + prepared[key] = value + else: + prepared[key] = value + return prepared if changed else params + + def _to_spanner_json_payload(value: Any) -> Any: """Prepare a value for Spanner JSON column parameter binding.""" if value is None: diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index ba18c4556..9058e1a68 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -350,7 +350,12 @@ def execute_partitioned_dml( if effective_request_options is not None: call_kwargs["request_options"] = effective_request_options - result = database.execute_partitioned_dml(sql, **call_kwargs) + result: Any = 0 + exc_handler = self.handle_database_exceptions() + with exc_handler: + result = database.execute_partitioned_dml(sql, **call_kwargs) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None return int(result) def execute( diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index 1996a971a..e8d803b4a 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -5,26 +5,15 @@ from typing_extensions import NotRequired -from sqlspec.adapters.spanner._typing import spanner_param_types as param_types -from sqlspec.adapters.spanner.type_converter import bytes_to_spanner, spanner_to_bytes +from sqlspec.adapters.spanner.type_converter import spanner_to_bytes from sqlspec.config import LitestarConfig +from sqlspec.core import TypedParameter from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ if TYPE_CHECKING: - from collections.abc import Callable - from typing import Protocol - - from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction from sqlspec.adapters.spanner.config import SpannerSyncConfig - class _DatabaseProtocol(Protocol): - def run_in_transaction(self, func: "Callable[[Transaction], Any]") -> Any: ... - - def update_ddl(self, ddl_statements: "list[str]") -> Any: ... - - def list_tables(self) -> Any: ... - __all__ = ("SpannerLitestarConfig", "SpannerSyncStore") @@ -90,9 +79,6 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) - def _database(self) -> "_DatabaseProtocol": - return cast("_DatabaseProtocol", self._config.get_database()) - def _datetime_to_timestamp(self, dt: "datetime | None") -> "datetime | None": if dt is None: return None @@ -110,23 +96,15 @@ def _timestamp_to_datetime(self, ts: "datetime | None") -> "datetime | None": def _build_params( self, key: str, expires_at: "datetime | None" = None, data: "bytes | None" = None ) -> "dict[str, Any]": - params: dict[str, Any] = {"session_id": key, "expires_at": self._datetime_to_timestamp(expires_at)} + ts = self._datetime_to_timestamp(expires_at) + params: dict[str, Any] = { + "session_id": key, + "expires_at": ts if ts is not None else TypedParameter(None, datetime), + } if data is not None: - params["data"] = bytes_to_spanner(data) + params["data"] = data return params - def _get_param_types( - self, session_id: bool = True, expires_at: bool = False, data: bool = False - ) -> "dict[str, Any]": - types: dict[str, Any] = {} - if session_id: - types["session_id"] = param_types.STRING - if expires_at: - types["expires_at"] = param_types.TIMESTAMP - if data: - types["data"] = param_types.BYTES - return types - def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": sql = f""" SELECT data, expires_at @@ -149,7 +127,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if result is None: return None - data = spanner_to_bytes(result.get("data")) + data = result.get("data") expires_at = self._timestamp_to_datetime(result.get("expires_at")) if renew_for is not None and expires_at is not None: diff --git a/sqlspec/adapters/spanner/type_converter.py b/sqlspec/adapters/spanner/type_converter.py index 207cb85d0..46f00f0cd 100644 --- a/sqlspec/adapters/spanner/type_converter.py +++ b/sqlspec/adapters/spanner/type_converter.py @@ -226,12 +226,15 @@ def coerce_params_for_spanner( Decimal: "NUMERIC", timedelta: "INTERVAL", UUID: "STRING", + dict: "JSON", "FLOAT32": "FLOAT32", "float32": "FLOAT32", "FLOAT64": "FLOAT64", "float64": "FLOAT64", "INTERVAL": "INTERVAL", "interval": "INTERVAL", + "JSON": "JSON", + "json": "JSON", } @@ -364,6 +367,8 @@ def _null_param_type(raw_value: Any, param_types: "SpannerParamTypesProtocol") - return getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) if resolver == "INTERVAL": return getattr(param_types, "INTERVAL", None) + if resolver == "JSON": + return _json_param_type() return getattr(param_types, resolver) if resolver is not None else None diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py index 112ffdb3d..ea91e6bc0 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -75,3 +75,88 @@ def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: expired_count = store._delete_expired() assert config.provide_session.call_count == 4 assert expired_count == 1 + + +def test_litestar_store_single_base64_roundtrip() -> None: + """Verify SpannerSyncStore passes raw bytes to driver.execute and decodes wire bytes once on _get.""" + from sqlspec.adapters.spanner.type_converter import bytes_to_spanner + from sqlspec.core import TypedParameter + + config = MagicMock(spec=SpannerSyncConfig) + config.extension_config = {"litestar": {"session_table": "sessions"}} + captured_params: list[dict[str, Any]] = [] + + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_result = MagicMock() + mock_result.rows_affected = 1 + + def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: + if isinstance(params, dict): + captured_params.append(params) + return mock_result + + mock_driver.execute.side_effect = mock_execute + mock_driver.select_one_or_none.return_value = {"data": bytes_to_spanner(b"raw-payload"), "expires_at": None} + config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) + + store = SpannerSyncStore(config=config) + store._set("session_1", b"raw-payload", expires_in=None) + + assert len(captured_params) == 1 + assert captured_params[0]["data"] == b"raw-payload" + assert isinstance(captured_params[0]["expires_at"], TypedParameter) + assert captured_params[0]["expires_at"].value is None + + fetched = store._get("session_1") + assert fetched == b"raw-payload" + + +def test_adk_memory_store_write_and_decode_json() -> None: + """Verify SpannerSyncADKMemoryStore prepares JSON/null write params and unwraps JsonObject.""" + from google.cloud.spanner_v1.data_types import JsonObject + + from sqlspec.adapters.spanner._typing import spanner_param_types as param_types + from sqlspec.adapters.spanner.adk import SpannerSyncADKMemoryStore + from sqlspec.core import TypedParameter + + config = MagicMock(spec=SpannerSyncConfig) + config.extension_config = {"adk": {"enable_memory": True}} + captured_params: list[dict[str, Any]] = [] + + mock_driver = MagicMock(spec=SpannerSyncDriver) + mock_result = MagicMock() + mock_result.rows_affected = 5 + + def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: + if isinstance(params, dict): + captured_params.append(params) + return mock_result + + mock_driver.execute.side_effect = mock_execute + config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) + + store = SpannerSyncADKMemoryStore(config=config) + store._run_write([ + ( + "INSERT INTO adk_memory VALUES (@content_json, @metadata_json, @owner_id)", + {"content_json": '{"text":"hi"}', "metadata_json": None, "owner_id": None}, + {"content_json": param_types.JSON, "metadata_json": param_types.JSON, "owner_id": param_types.STRING}, + ) + ]) + + assert len(captured_params) == 1 + assert captured_params[0]["content_json"] == {"text": "hi"} + assert isinstance(captured_params[0]["metadata_json"], TypedParameter) + assert captured_params[0]["metadata_json"].original_type is dict + assert isinstance(captured_params[0]["owner_id"], TypedParameter) + assert captured_params[0]["owner_id"].original_type is str + + deleted = store._execute_update( + "DELETE FROM adk_memory WHERE session_id = @session_id", + {"session_id": "s1"}, + {"session_id": param_types.STRING}, + ) + assert deleted == 5 + + decoded = store._decode_json(JsonObject({"k": "v"})) + assert decoded == {"k": "v"} diff --git a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py index 9ef0df154..a68be300e 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py +++ b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py @@ -40,3 +40,11 @@ def test_null_timedelta_param_types() -> None: types = infer_spanner_param_types(params) assert "duration" in types assert types["duration"].code == TypeCode.INTERVAL + + +def test_null_json_param_types() -> None: + """Verify that null TypedParameter with dict or JSON resolves to JSON.""" + params = {"meta": TypedParameter(None, dict), "payload": TypedParameter(None, "JSON")} + types = infer_spanner_param_types(params) + assert types["meta"].code == TypeCode.JSON + assert types["payload"].code == TypeCode.JSON From 30b2ef9c87b9214feb8fb8134634ee1efc98c1e1 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 02:06:17 +0000 Subject: [PATCH 07/15] test(spanner): use typed spanner_json helper in test_spanner_stores --- tests/unit/adapters/test_spanner/test_spanner_stores.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py index ea91e6bc0..eb5259ca3 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -113,10 +113,9 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: def test_adk_memory_store_write_and_decode_json() -> None: """Verify SpannerSyncADKMemoryStore prepares JSON/null write params and unwraps JsonObject.""" - from google.cloud.spanner_v1.data_types import JsonObject - from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.adapters.spanner.adk import SpannerSyncADKMemoryStore + from sqlspec.adapters.spanner.type_converter import spanner_json from sqlspec.core import TypedParameter config = MagicMock(spec=SpannerSyncConfig) @@ -158,5 +157,5 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: ) assert deleted == 5 - decoded = store._decode_json(JsonObject({"k": "v"})) + decoded = store._decode_json(spanner_json({"k": "v"})) assert decoded == {"k": "v"} From 6ac958b1e0ca98f3f81ce1bf9663d58c6b53c80c Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 15:38:06 +0000 Subject: [PATCH 08/15] fix(spanner): use transaction=True in ADK and Litestar store writes and fix mutation transaction exit --- sqlspec/adapters/bigquery/core.py | 6 +- sqlspec/adapters/spanner/adk/store.py | 12 +-- sqlspec/adapters/spanner/config.py | 83 +++++++------------ sqlspec/adapters/spanner/driver.py | 44 +++++++--- sqlspec/adapters/spanner/litestar/store.py | 10 +-- sqlspec/core/parameters/_types.py | 2 +- .../test_spanner/test_batch_write_api.py | 7 +- .../unit/adapters/test_spanner/test_config.py | 75 +++++++++++++++++ .../test_spanner/test_litestar_store.py | 52 ++++++++---- .../test_load_from_arrow_mutations.py | 3 +- .../test_spanner_partitioned_dml.py | 19 ++--- .../test_spanner/test_spanner_stores.py | 27 +++--- 12 files changed, 213 insertions(+), 127 deletions(-) diff --git a/sqlspec/adapters/bigquery/core.py b/sqlspec/adapters/bigquery/core.py index c7d323eff..1ea8205f1 100644 --- a/sqlspec/adapters/bigquery/core.py +++ b/sqlspec/adapters/bigquery/core.py @@ -254,7 +254,7 @@ def create_parameters(parameters: Any, json_serializer: "Callable[[Any], str] | if _is_query_parameter(value): bq_parameters.append(cast("BigQueryParam", value)) continue - declared_type: type[Any] | str | None = None + declared_type: type[Any] | None = None if type(value) is TypedParameter: declared_type = value.original_type actual_value = value.value @@ -1007,9 +1007,7 @@ def _load_bigquery_module() -> Any: return _BIGQUERY_MODULE -def _query_parameter_type( - value: Any, declared_type: "type[Any] | str | None" = None -) -> "tuple[str | None, str | None]": +def _query_parameter_type(value: Any, declared_type: "type[Any] | None" = None) -> "tuple[str | None, str | None]": """Determine BigQuery parameter type from Python value. Args: diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index c5f7ef00b..e176214b8 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -225,7 +225,7 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: for sql, params, types in statements: driver.execute(sql, _prepare_spanner_write_params(params, types)) @@ -634,7 +634,7 @@ def _delete_expired_events(self, before: datetime, app_name: "str | None" = None if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(sql, params) return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) @@ -644,7 +644,7 @@ def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(sql, params) return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) @@ -654,7 +654,7 @@ def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | No if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(sql, params) return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) @@ -900,12 +900,12 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: for sql, params, types in statements: driver.execute(sql, _prepare_spanner_write_params(params, types)) def _execute_update(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> int: - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(sql, _prepare_spanner_write_params(params, types)) return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index a678205e9..ae2ac9710 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -5,6 +5,7 @@ from typing_extensions import NotRequired +from sqlspec.adapters.spanner import _typing as spanner_typing from sqlspec.adapters.spanner._typing import SpannerConnection from sqlspec.adapters.spanner._typing import SpannerTransactionType as TransactionType from sqlspec.adapters.spanner.core import apply_driver_features, default_statement_config @@ -33,6 +34,7 @@ from sqlspec.adapters.spanner._typing import SpannerDirectedReadOptions as DirectedReadOptions from sqlspec.adapters.spanner._typing import SpannerEncryptionConfig as EncryptionConfig from sqlspec.adapters.spanner._typing import SpannerExecuteSqlRequest as ExecuteSqlRequest + from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool from sqlspec.adapters.spanner._typing import SpannerRequestOptions as RequestOptions from sqlspec.adapters.spanner._typing import SpannerRetry as Retry from sqlspec.config import ExtensionConfigs @@ -215,23 +217,18 @@ def __init__(self, config: "SpannerSyncConfig", transaction: bool = False) -> No def __enter__(self) -> SpannerConnection: database = self._config.get_database() if self._transaction: - manager = getattr(database, "sessions_manager", None) - if manager is not None and hasattr(manager, "get_session"): - self._session = manager.get_session(TransactionType.READ_WRITE) - try: - txn = self._session.transaction() - txn.__enter__() - self._connection = cast("SpannerConnection", txn) - except Exception: - manager.put_session(self._session) - self._session = None - raise - else: - return self._connection - txn = database.transaction() - self._session = txn - self._connection = cast("SpannerConnection", txn.__enter__()) - return self._connection + manager = cast("Any", database).sessions_manager + self._session = manager.get_session(TransactionType.READ_WRITE) + try: + txn = self._session.transaction() + txn.__enter__() + self._connection = cast("SpannerConnection", txn) + except Exception: + manager.put_session(self._session) + self._session = None + raise + else: + return self._connection self._session = cast("Any", database).snapshot(multi_use=True) self._connection = cast("SpannerConnection", self._session.__enter__()) return self._connection @@ -242,33 +239,18 @@ def __exit__( if self._transaction and self._connection: txn = cast("Any", self._connection) try: + rolled_back = bool(getattr(txn, "rolled_back", False)) + committed = getattr(txn, "committed", None) + txn_id = getattr(txn, "_transaction_id", None) + mutations = cast("list[Any] | None", getattr(txn, "_mutations", None)) if exc_type is None: - try: - txn_id = txn._transaction_id - except AttributeError: - txn_id = None - mutations = cast("list[Any] | None", getattr(txn, "_mutations", None)) - try: - committed = txn.committed - except AttributeError: - committed = None - if committed is None and (txn_id is not None or bool(mutations)): + if not rolled_back and committed is None and (txn_id is not None or bool(mutations)): txn.commit() - else: - try: - rollback_txn_id = txn._transaction_id - except AttributeError: - rollback_txn_id = None - if rollback_txn_id is not None: - txn.rollback() + elif not rolled_back and committed is None and txn_id is not None: + txn.rollback() finally: if self._session: - db = self._config.get_database() - manager = getattr(db, "sessions_manager", None) - if manager is not None and hasattr(manager, "put_session"): - manager.put_session(self._session) - elif hasattr(self._session, "__exit__"): - self._session.__exit__(exc_type, exc_val, exc_tb) + cast("Any", self._config.get_database()).sessions_manager.put_session(self._session) elif self._session: self._session.__exit__(exc_type, exc_val, exc_tb) @@ -341,9 +323,7 @@ def __init__( if enable_multiplexed and "pool_type" not in self.connection_config: self.connection_config["pool_type"] = None elif not enable_multiplexed and "pool_type" not in self.connection_config: - from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool - - self.connection_config["pool_type"] = PingingPool + self.connection_config["pool_type"] = spanner_typing.SpannerPingingPool self.connection_config.setdefault("ping_interval", 1800) statement_config = statement_config or default_statement_config @@ -365,11 +345,9 @@ def __init__( self._database: Database | None = None def _get_client(self) -> "Client": - from sqlspec.adapters.spanner._typing import SpannerClient as Client - if self._client is None: client_kwargs = self._connection_kwargs_for(_CLIENT_CONFIG_FIELDS) - self._client = Client(**client_kwargs) + self._client = spanner_typing.SpannerClient(**client_kwargs) return self._client def get_database(self) -> "Database": @@ -413,11 +391,6 @@ def create_connection(self) -> SpannerConnection: return cast("SpannerConnection", self.get_database().snapshot(multi_use=True)) # type: ignore[no-untyped-call] def _create_pool(self) -> "AbstractSessionPool": - from sqlspec.adapters.spanner._typing import SpannerAbstractSessionPool as AbstractSessionPool - from sqlspec.adapters.spanner._typing import SpannerBurstyPool as BurstyPool - from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool - from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool - instance_id = self.connection_config.get("instance_id") database_id = self.connection_config.get("database_id") if not instance_id or not database_id: @@ -427,18 +400,18 @@ def _create_pool(self) -> "AbstractSessionPool": raw_pool_type = self.connection_config.get("pool_type") pool_type: type[AbstractSessionPool | PingingPool] if raw_pool_type is None or raw_pool_type == "multiplexed": - pool_type = PingingPool + pool_type = spanner_typing.SpannerPingingPool else: pool_type = cast("type[AbstractSessionPool]", raw_pool_type) labels = self.connection_config.get("session_labels", self.connection_config.get("labels")) pool_kwargs: dict[str, Any] = self._pool_base_kwargs(labels=cast("dict[str, str] | None", labels)) - if issubclass(pool_type, PingingPool): + if issubclass(pool_type, spanner_typing.SpannerPingingPool): self.connection_config.setdefault("ping_interval", 1800) pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "ping_interval"})) - elif issubclass(pool_type, FixedSizePool): + elif issubclass(pool_type, spanner_typing.SpannerFixedSizePool): pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "max_age_minutes"})) - elif issubclass(pool_type, BurstyPool): + elif issubclass(pool_type, spanner_typing.SpannerBurstyPool): target_size = self.connection_config.get("target_size", self.connection_config.get("size")) if target_size is not None: pool_kwargs["target_size"] = target_size diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 9058e1a68..2962a87d3 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -239,18 +239,36 @@ def begin(self) -> None: return None def commit(self) -> None: - if isinstance(self.connection, SpannerTransaction) or supports_write(self.connection): - writer = cast("_SpannerWriteProtocol", self.connection) - if getattr(writer, "committed", None) is not None: - return - if callable(getattr(writer, "commit", None)): - writer.commit() + writer = cast("_SpannerWriteProtocol", self.connection) + if not ( + isinstance(self.connection, SpannerTransaction) + or supports_write(self.connection) + or callable(getattr(writer, "commit", None)) + ): + return + if getattr(writer, "rolled_back", False) or getattr(writer, "committed", None) is not None: + return + if ( + getattr(writer, "_transaction_id", None) is not None + or bool(getattr(writer, "_mutations", None)) + or (not hasattr(writer, "_transaction_id") and not hasattr(writer, "_mutations")) + ) and callable(getattr(writer, "commit", None)): + writer.commit() def rollback(self) -> None: - if isinstance(self.connection, SpannerTransaction) or supports_write(self.connection): - writer = cast("_SpannerWriteProtocol", self.connection) - if callable(getattr(writer, "rollback", None)): - writer.rollback() + writer = cast("_SpannerWriteProtocol", self.connection) + if not ( + isinstance(self.connection, SpannerTransaction) + or supports_write(self.connection) + or callable(getattr(writer, "rollback", None)) + ): + return + if getattr(writer, "rolled_back", False) or getattr(writer, "committed", None) is not None: + return + if (getattr(writer, "_transaction_id", None) is not None or not hasattr(writer, "_transaction_id")) and ( + callable(getattr(writer, "rollback", None)) + ): + writer.rollback() def create_savepoint(self, name: str) -> None: """Raise because Spanner does not support savepoints. @@ -298,7 +316,7 @@ def _get_database(self) -> Any: return database return None - def execute_partitioned_dml( + def _execute_partitioned_dml( self, statement: "SQL | str", *parameters: Any, @@ -516,12 +534,12 @@ def load_from_arrow( arrow_table = self._coerce_arrow_table(source) if overwrite: - self.execute_partitioned_dml(f"DELETE FROM {table} WHERE TRUE") + self._execute_partitioned_dml(f"DELETE FROM {table} WHERE TRUE") columns, records = self._arrow_table_to_rows(arrow_table) if records: chunks = self._chunk_mutation_rows(columns, records) - if self.driver_features.get("enable_batch_write_api") and not overwrite: + if self.driver_features.get("enable_batch_write_api"): self._batch_write_mutations(table, columns, chunks) else: conn = self.connection diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index e8d803b4a..ab6243a99 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -140,7 +140,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if self._shard_count > 1: update_sql = f"{update_sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = self._build_params(key, new_expires) - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: driver.execute(update_sql, params) return spanner_to_bytes(data) @@ -164,7 +164,7 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No VALUES (@session_id, @data, @expires_at, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()) """ - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(update_sql, params) rows_affected = getattr(result, "rows_affected", None) has_rows = rows_affected > 0 if isinstance(rows_affected, int) else bool(getattr(result, "rowcount", None)) @@ -176,12 +176,12 @@ def _delete(self, key: str) -> None: if self._shard_count > 1: sql = f"{sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = {"session_id": key} - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: driver.execute(sql, params) def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name} WHERE TRUE" - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: driver.execute(sql) def _exists(self, key: str) -> bool: @@ -217,7 +217,7 @@ def _delete_expired(self) -> int: DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= CURRENT_TIMESTAMP() """ - with self._config.provide_session() as driver: + with self._config.provide_session(transaction=True) as driver: result = driver.execute(sql) rows_affected = getattr(result, "rows_affected", None) if isinstance(rows_affected, int): diff --git a/sqlspec/core/parameters/_types.py b/sqlspec/core/parameters/_types.py index e7fd25b54..f56491a60 100644 --- a/sqlspec/core/parameters/_types.py +++ b/sqlspec/core/parameters/_types.py @@ -180,7 +180,7 @@ def __init__( self, value: Any, original_type: "type | str | None" = None, semantic_name: "str | None" = None ) -> None: self.value = value - self.original_type = original_type or type(value) + self.original_type: Any = original_type or type(value) self.semantic_name = semantic_name self._hash: int | None = None diff --git a/tests/unit/adapters/test_spanner/test_batch_write_api.py b/tests/unit/adapters/test_spanner/test_batch_write_api.py index 30df6b68f..198225c25 100644 --- a/tests/unit/adapters/test_spanner/test_batch_write_api.py +++ b/tests/unit/adapters/test_spanner/test_batch_write_api.py @@ -115,12 +115,13 @@ def test_batch_write_splits_before_crossing_mutation_group_cell_cap(batch_write_ assert [len(chunk) for chunk in chunks] == [26_666, 1] -def test_batch_write_overwrite_uses_transactional_mutations(batch_write_driver: SpannerSyncDriver) -> None: +def test_batch_write_overwrite_uses_batch_write_after_partitioned_dml(batch_write_driver: SpannerSyncDriver) -> None: conn = cast("_FakeBatchTransaction", batch_write_driver.connection) batch_write_driver.load_from_arrow("users", pa.table({"id": [1]}), overwrite=True) assert ( conn.database.partitioned_dml_calls and "DELETE FROM users WHERE TRUE" in conn.database.partitioned_dml_calls[0] ) - assert conn.insert_or_update_calls == [("users", ["id"], [[1]])] - assert conn.database.mutation_groups_obj.batch_write_calls == 0 + assert conn.insert_or_update_calls == [] + assert conn.database.mutation_groups_obj.batch_write_calls == 1 + assert conn.database.mutation_groups_obj.groups[0].calls == [("users", ["id"], [[1]])] diff --git a/tests/unit/adapters/test_spanner/test_config.py b/tests/unit/adapters/test_spanner/test_config.py index 3fa465353..d8ae400c8 100644 --- a/tests/unit/adapters/test_spanner/test_config.py +++ b/tests/unit/adapters/test_spanner/test_config.py @@ -527,6 +527,81 @@ def __init__(self) -> None: assert db.session_obj.txn.rollback_calls == 0 +def test_transaction_context_and_driver_respect_rolled_back_state() -> None: + """Explicit driver.rollback() must prevent auto-commit or duplicate rollback on exit.""" + + class _Txn: + def __init__(self) -> None: + self._transaction_id: str | None = "txn-1" + self._mutations: list[object] = [object()] + self.committed: object | None = None + self.rolled_back = False + self.commit_calls = 0 + self.rollback_calls = 0 + + def __enter__(self): + return self + + def __exit__(self, *_: object) -> None: + return None + + def commit(self) -> None: + self.commit_calls += 1 + self.committed = object() + + def rollback(self) -> None: + if self.rolled_back: + msg = "Transaction already rolled back." + raise ValueError(msg) + self.rollback_calls += 1 + self.rolled_back = True + + class _Session: + def __init__(self) -> None: + self.txn = _Txn() + + def transaction(self) -> _Txn: + return self.txn + + class _SessionsManager: + def __init__(self, session: _Session) -> None: + self.session = session + self.returned = 0 + + def get_session(self, _transaction_type: object) -> _Session: + return self.session + + def put_session(self, _session: object) -> None: + self.returned += 1 + + class _DB: + def __init__(self) -> None: + self.session_obj = _Session() + self.sessions_manager = _SessionsManager(self.session_obj) + + db = _DB() + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + setattr(config, "get_database", lambda: db) + + with config.provide_session(transaction=True) as driver: + driver.rollback() + driver.rollback() + driver.commit() + + assert db.session_obj.txn.rollback_calls == 1 + assert db.session_obj.txn.commit_calls == 0 + + db_err = _DB() + setattr(config, "get_database", lambda: db_err) + with pytest.raises(RuntimeError, match="abort"), config.provide_session(transaction=True) as driver: + driver.rollback() + msg = "abort" + raise RuntimeError(msg) + + assert db_err.session_obj.txn.rollback_calls == 1 + assert db_err.session_obj.txn.commit_calls == 0 + + def test_provide_session_uses_batch_when_transaction_requested() -> None: """Driver should receive transaction connection when transaction=True.""" diff --git a/tests/unit/adapters/test_spanner/test_litestar_store.py b/tests/unit/adapters/test_spanner/test_litestar_store.py index 4aff468fb..ad2432f55 100644 --- a/tests/unit/adapters/test_spanner/test_litestar_store.py +++ b/tests/unit/adapters/test_spanner/test_litestar_store.py @@ -1,11 +1,13 @@ +from datetime import datetime, timezone from typing import Any -from unittest.mock import MagicMock +from unittest.mock import MagicMock, call from sqlspec.adapters.spanner.litestar import SpannerSyncStore +from sqlspec.adapters.spanner.type_converter import bytes_to_spanner def test_set_uses_session() -> None: - """Verify _set uses config.provide_session for write operations.""" + """Verify _set uses config.provide_session(transaction=True) for write operations.""" driver = MagicMock() driver.execute.return_value = MagicMock(rows_affected=1) cm = _context_manager_yielding(driver) @@ -15,13 +17,13 @@ def test_set_uses_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - store._set("s1", b"data", None) # pyright: ignore + store._set("s1", b"data", None) - config.provide_session.assert_called_once() + config.provide_session.assert_called_once_with(transaction=True) def test_delete_uses_session() -> None: - """Verify _delete uses config.provide_session for write operations.""" + """Verify _delete uses config.provide_session(transaction=True) for write operations.""" driver = MagicMock() cm = _context_manager_yielding(driver) @@ -30,13 +32,13 @@ def test_delete_uses_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - store._delete("s1") # pyright: ignore + store._delete("s1") - config.provide_session.assert_called_once() + config.provide_session.assert_called_once_with(transaction=True) def test_delete_all_uses_session() -> None: - """Verify _delete_all uses config.provide_session for write operations.""" + """Verify _delete_all uses config.provide_session(transaction=True) for write operations.""" driver = MagicMock() cm = _context_manager_yielding(driver) @@ -45,13 +47,13 @@ def test_delete_all_uses_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - store._delete_all() # pyright: ignore + store._delete_all() - config.provide_session.assert_called_once() + config.provide_session.assert_called_once_with(transaction=True) def test_delete_expired_uses_session() -> None: - """Verify _delete_expired uses config.provide_session for write operations.""" + """Verify _delete_expired uses config.provide_session(transaction=True) for write operations.""" driver = MagicMock() driver.execute.return_value = MagicMock(rows_affected=3) cm = _context_manager_yielding(driver) @@ -61,9 +63,9 @@ def test_delete_expired_uses_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - result = store._delete_expired() # pyright: ignore + result = store._delete_expired() - config.provide_session.assert_called_once() + config.provide_session.assert_called_once_with(transaction=True) assert result == 3 @@ -89,12 +91,32 @@ def test_get_uses_snapshot_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - result = store._get("s1") # pyright: ignore + result = store._get("s1") config.provide_session.assert_called_once_with() assert result is None +def test_get_renewal_uses_transaction_session() -> None: + """Verify _get token renewal uses provide_session(transaction=True) for the UPDATE.""" + driver = MagicMock() + driver.select_one_or_none.return_value = { + "data": bytes_to_spanner(b"val"), + "expires_at": datetime(2030, 1, 1, tzinfo=timezone.utc), + } + cm = _context_manager_yielding(driver) + + config = MagicMock() + config.extension_config = {"litestar": {"session_table": "sess"}} + config.provide_session.return_value = cm + + store = SpannerSyncStore(config) + result = store._get("s1", renew_for=60) + + assert result == b"val" + assert config.provide_session.call_args_list == [call(), call(transaction=True)] + + def test_exists_uses_snapshot_session() -> None: """Verify _exists uses snapshot session for read operations.""" driver = MagicMock() @@ -106,7 +128,7 @@ def test_exists_uses_snapshot_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - result = store._exists("s1") # pyright: ignore + result = store._exists("s1") config.provide_session.assert_called_once_with() assert result is True diff --git a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py index fbd76fb90..1969ee7d8 100644 --- a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py +++ b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py @@ -1,6 +1,7 @@ """Spanner load_from_arrow mutations transport (insert_or_update).""" from typing import Any, cast +from unittest.mock import MagicMock import pyarrow as pa import pytest @@ -22,8 +23,6 @@ def __init__(self) -> None: self.insert_or_update_calls: list[tuple[str, list[str], list[list[Any]]]] = [] self.execute_update_calls: list[str] = [] self.committed = None - from unittest.mock import MagicMock - self._database = MagicMock() self._database.execute_partitioned_dml.return_value = 0 diff --git a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py index 4ebb55cd2..4f905bf6a 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py +++ b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py @@ -2,14 +2,17 @@ from unittest.mock import MagicMock +import pytest from google.cloud.spanner_v1.types.type import TypeCode from sqlspec.adapters.spanner.core import default_statement_config from sqlspec.adapters.spanner.driver import SpannerSyncDriver +from sqlspec.core import SQL +from sqlspec.exceptions import SQLConversionError def test_driver_execute_partitioned_dml() -> None: - """Verify that driver.execute_partitioned_dml delegates to database.execute_partitioned_dml.""" + """Verify that driver._execute_partitioned_dml delegates to database.execute_partitioned_dml.""" mock_db = MagicMock() mock_db.execute_partitioned_dml.return_value = 42 @@ -20,7 +23,7 @@ def test_driver_execute_partitioned_dml() -> None: connection=mock_connection, statement_config=default_statement_config, driver_features={} ) - rows = driver.execute_partitioned_dml("DELETE FROM large_table WHERE active = FALSE") + rows = driver._execute_partitioned_dml("DELETE FROM large_table WHERE active = FALSE") assert rows == 42 mock_db.execute_partitioned_dml.assert_called_once() sql = mock_db.execute_partitioned_dml.call_args[0][0] @@ -39,7 +42,7 @@ def test_driver_execute_partitioned_dml_with_parameters() -> None: connection=mock_connection, statement_config=default_statement_config, driver_features={} ) - rows = driver.execute_partitioned_dml( + rows = driver._execute_partitioned_dml( "UPDATE large_table SET status = :status WHERE threshold > :limit", {"status": "archived", "limit": 100} ) assert rows == 10 @@ -54,8 +57,6 @@ def test_driver_execute_partitioned_dml_with_parameters() -> None: def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: """Verify executing partitioned DML with SQL object, query options, and request options.""" - from sqlspec.core import SQL - mock_db = MagicMock() mock_db.execute_partitioned_dml.return_value = 50 @@ -70,7 +71,7 @@ def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: mock_query_options = MagicMock() mock_request_options = MagicMock() - rows = driver.execute_partitioned_dml( + rows = driver._execute_partitioned_dml( statement, query_options=mock_query_options, request_options=mock_request_options, @@ -86,10 +87,6 @@ def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: def test_driver_execute_partitioned_dml_no_database_raises() -> None: """Verify error raised when database cannot be resolved.""" - import pytest - - from sqlspec.exceptions import SQLConversionError - mock_connection = MagicMock() mock_connection._session = None mock_connection._database = None @@ -99,4 +96,4 @@ def test_driver_execute_partitioned_dml_no_database_raises() -> None: ) with pytest.raises(SQLConversionError, match="Could not resolve Spanner database"): - driver.execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") + driver._execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py index eb5259ca3..63d960441 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ b/tests/unit/adapters/test_spanner/test_spanner_stores.py @@ -3,10 +3,13 @@ from typing import Any from unittest.mock import MagicMock -from sqlspec.adapters.spanner.adk import SpannerSyncADKStore +from sqlspec.adapters.spanner._typing import spanner_param_types as param_types +from sqlspec.adapters.spanner.adk import SpannerSyncADKMemoryStore, SpannerSyncADKStore from sqlspec.adapters.spanner.config import SpannerSyncConfig from sqlspec.adapters.spanner.driver import SpannerSyncDriver from sqlspec.adapters.spanner.litestar import SpannerSyncStore +from sqlspec.adapters.spanner.type_converter import bytes_to_spanner, spanner_json +from sqlspec.core import TypedParameter def _context_manager_yielding(value: Any) -> Any: @@ -21,7 +24,7 @@ def __exit__(self, *_: Any) -> None: def test_adk_store_run_write_routes_through_provide_session() -> None: - """Verify that SpannerSyncADKStore._run_write executes via config.provide_session.""" + """Verify that SpannerSyncADKStore._run_write executes via config.provide_session(transaction=True).""" config = MagicMock(spec=SpannerSyncConfig) executed_statements: list[tuple[str, Any]] = [] @@ -39,12 +42,12 @@ def test_adk_store_run_write_routes_through_provide_session() -> None: ] store._run_write(statements) - config.provide_session.assert_called_once() + config.provide_session.assert_called_once_with(transaction=True) assert len(executed_statements) == 2 def test_litestar_store_writes_route_through_provide_session() -> None: - """Verify that SpannerSyncStore write operations execute via config.provide_session.""" + """Verify that SpannerSyncStore write operations execute via config.provide_session(transaction=True).""" config = MagicMock(spec=SpannerSyncConfig) config.extension_config = {"litestar": {"session_table": "sessions"}} executed_sqls: list[str] = [] @@ -65,23 +68,24 @@ def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: store._set("session_1", b"payload", expires_in=3600) assert config.provide_session.call_count == 1 + config.provide_session.assert_called_with(transaction=True) store._delete("session_1") assert config.provide_session.call_count == 2 + config.provide_session.assert_called_with(transaction=True) store._delete_all() assert config.provide_session.call_count == 3 + config.provide_session.assert_called_with(transaction=True) expired_count = store._delete_expired() assert config.provide_session.call_count == 4 + config.provide_session.assert_called_with(transaction=True) assert expired_count == 1 def test_litestar_store_single_base64_roundtrip() -> None: """Verify SpannerSyncStore passes raw bytes to driver.execute and decodes wire bytes once on _get.""" - from sqlspec.adapters.spanner.type_converter import bytes_to_spanner - from sqlspec.core import TypedParameter - config = MagicMock(spec=SpannerSyncConfig) config.extension_config = {"litestar": {"session_table": "sessions"}} captured_params: list[dict[str, Any]] = [] @@ -101,6 +105,7 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: store = SpannerSyncStore(config=config) store._set("session_1", b"raw-payload", expires_in=None) + config.provide_session.assert_called_once_with(transaction=True) assert len(captured_params) == 1 assert captured_params[0]["data"] == b"raw-payload" @@ -113,11 +118,6 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: def test_adk_memory_store_write_and_decode_json() -> None: """Verify SpannerSyncADKMemoryStore prepares JSON/null write params and unwraps JsonObject.""" - from sqlspec.adapters.spanner._typing import spanner_param_types as param_types - from sqlspec.adapters.spanner.adk import SpannerSyncADKMemoryStore - from sqlspec.adapters.spanner.type_converter import spanner_json - from sqlspec.core import TypedParameter - config = MagicMock(spec=SpannerSyncConfig) config.extension_config = {"adk": {"enable_memory": True}} captured_params: list[dict[str, Any]] = [] @@ -142,6 +142,7 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: {"content_json": param_types.JSON, "metadata_json": param_types.JSON, "owner_id": param_types.STRING}, ) ]) + config.provide_session.assert_called_once_with(transaction=True) assert len(captured_params) == 1 assert captured_params[0]["content_json"] == {"text": "hi"} @@ -155,6 +156,8 @@ def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: {"session_id": "s1"}, {"session_id": param_types.STRING}, ) + assert config.provide_session.call_count == 2 + config.provide_session.assert_called_with(transaction=True) assert deleted == 5 decoded = store._decode_json(spanner_json({"k": "v"})) From 4ce62d6b323872eb9e8c00d1dfd2eeca24263bdc Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 20:54:50 +0000 Subject: [PATCH 09/15] fix(spanner): preserve native transactions and narrow optimization scope --- sqlspec/adapters/spanner/adk/store.py | 199 ++++++++---------- sqlspec/adapters/spanner/config.py | 96 ++++----- sqlspec/adapters/spanner/core.py | 9 +- sqlspec/adapters/spanner/driver.py | 175 +++------------ sqlspec/adapters/spanner/litestar/store.py | 107 +++++++--- sqlspec/adapters/spanner/type_converter.py | 68 +----- sqlspec/core/parameters/_types.py | 9 +- sqlspec/protocols.py | 2 - .../adapters/_shared/_driver_type_system.py | 1 - .../test_spanner/test_batch_write_api.py | 16 +- .../unit/adapters/test_spanner/test_config.py | 75 ------- .../test_spanner/test_litestar_store.py | 90 +++----- .../test_load_from_arrow_mutations.py | 8 +- .../test_spanner_arrow_overwrite.py | 35 --- .../test_spanner/test_spanner_batch_write.py | 59 ------ .../test_spanner_last_statement.py | 58 ----- .../test_spanner_partitioned_dml.py | 99 --------- .../test_spanner/test_spanner_pinging_pool.py | 51 ----- .../test_spanner/test_spanner_pool.py | 71 ------- .../test_spanner_query_options.py | 101 --------- .../test_spanner_request_options.py | 105 --------- .../test_spanner/test_spanner_stores.py | 164 --------------- .../test_spanner_type_inference.py | 24 +-- .../test_spanner/test_spanner_vector.py | 48 ----- 24 files changed, 296 insertions(+), 1374 deletions(-) delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_batch_write.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_last_statement.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_pool.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_query_options.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_request_options.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_stores.py delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_vector.py diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index e176214b8..5facfb347 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -9,9 +9,8 @@ from sqlspec.adapters.spanner._typing import SpannerNotFound as NotFound from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.adapters.spanner.config import SpannerSyncConfig -from sqlspec.adapters.spanner.core import _unwrap_spanner_json_object from sqlspec.config import ADKConfig -from sqlspec.core import TypedParameter +from sqlspec.exceptions import OperationalError from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore from sqlspec.protocols import SpannerParamTypesProtocol @@ -21,6 +20,7 @@ from collections.abc import Sequence from sqlspec.adapters.spanner._typing import SpannerDatabase as Database + from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction from sqlspec.extensions.adk import SessionOrderBy, StoredMemory __all__ = ("SpannerADKConfig", "SpannerADKRetentionConfig", "SpannerSyncADKMemoryStore", "SpannerSyncADKStore") @@ -225,9 +225,7 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - with self._config.provide_session(transaction=True) as driver: - for sql, params, types in statements: - driver.execute(sql, _prepare_spanner_write_params(params, types)) + self._database().run_in_transaction(_SpannerWriteJob(statements)) # type: ignore[no-untyped-call] def _session_param_types(self, include_owner: bool) -> "dict[str, Any]": json_type = _json_param_type() @@ -267,30 +265,22 @@ def _metadata_param_types(self) -> "dict[str, Any]": return {"key": SPANNER_PARAM_TYPES.STRING, "value": SPANNER_PARAM_TYPES.STRING} def _decode_state(self, raw: Any) -> Any: - if raw is None: - return None if isinstance(raw, str): - try: - return from_json(raw) - except Exception: - return raw - return _unwrap_spanner_json_object(raw) + return from_json(raw) + return raw def _decode_json(self, raw: Any) -> Any: if raw is None: return None if isinstance(raw, str): - try: - return from_json(raw) - except Exception: - return raw - return _unwrap_spanner_json_object(raw) + return from_json(raw) + return raw def _create_session( self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None ) -> StoredSession: - state_payload = _to_spanner_json_payload(state) - params: dict[str, Any] = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_payload} + state_json = to_json(state) + params: dict[str, Any] = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_json} columns = "id, app_name, user_id, state, create_time, update_time" values = "@id, @app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()" if self._owner_id_column_name: @@ -372,7 +362,7 @@ def _get_session( return record def _update_session_state(self, app_name: str, user_id: str, session_id: str, state: "dict[str, Any]") -> None: - params = {"app_name": app_name, "user_id": user_id, "id": session_id, "state": _to_spanner_json_payload(state)} + params = {"app_name": app_name, "user_id": user_id, "id": session_id, "state": to_json(state)} json_type = _json_param_type() sql = f""" UPDATE {self._session_table} @@ -450,18 +440,13 @@ def _delete_session(self, app_name: str, user_id: str, session_id: str) -> None: ) delete_events_sql = f"DELETE FROM {self._events_table} WHERE session_id = @session_id{shard_clause}" delete_session_sql = f"DELETE FROM {self._session_table} WHERE app_name = @app_name AND user_id = @user_id AND id = @session_id{shard_clause}" - delete_events_params = {"session_id": session_id} - delete_events_types = {"session_id": SPANNER_PARAM_TYPES.STRING} - delete_session_params = {"app_name": app_name, "user_id": user_id, "session_id": session_id} - delete_session_types = { + params = {"app_name": app_name, "user_id": user_id, "session_id": session_id} + types = { "app_name": SPANNER_PARAM_TYPES.STRING, "user_id": SPANNER_PARAM_TYPES.STRING, "session_id": SPANNER_PARAM_TYPES.STRING, } - self._run_write([ - (delete_events_sql, delete_events_params, delete_events_types), - (delete_session_sql, delete_session_params, delete_session_types), - ]) + self._run_write([(delete_events_sql, params, types), (delete_session_sql, params, types)]) def _append_event_and_update_state( self, @@ -498,7 +483,7 @@ def _append_event_and_update_state( "session_id": event_record["session_id"], "invocation_id": event_record["invocation_id"], "timestamp": event_record["timestamp"], - "event_data": _to_spanner_json_payload(event_record["event_data"]), + "event_data": to_json(event_record["event_data"]), } insert_sql = f""" INSERT INTO {self._events_table} (id, app_name, user_id, session_id, invocation_id, timestamp, event_data) @@ -510,7 +495,7 @@ def _append_event_and_update_state( "app_name": app_name, "user_id": user_id, "id": session_id, - "state": _to_spanner_json_payload(state), + "state": to_json(state), } update_sql = f""" UPDATE {self._session_table} @@ -539,7 +524,7 @@ def _append_event_and_update_state( INSERT OR UPDATE {self._app_state_table} (app_name, state, update_time) VALUES (@app_name, @state, PENDING_COMMIT_TIMESTAMP()) """, - {"app_name": app_name, "state": _to_spanner_json_payload(app_state)}, + {"app_name": app_name, "state": to_json(app_state)}, self._app_state_param_types(), )) if user_state is not None: @@ -548,7 +533,7 @@ def _append_event_and_update_state( INSERT OR UPDATE {self._user_state_table} (app_name, user_id, state, update_time) VALUES (@app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP()) """, - {"app_name": app_name, "user_id": user_id, "state": _to_spanner_json_payload(user_state)}, + {"app_name": app_name, "user_id": user_id, "state": to_json(user_state)}, self._user_state_param_types(), )) @@ -568,7 +553,7 @@ def _insert_event(self, event_record: "StoredEvent") -> None: "session_id": event_record["session_id"], "invocation_id": event_record["invocation_id"], "timestamp": event_record["timestamp"], - "event_data": _to_spanner_json_payload(event_record["event_data"]), + "event_data": to_json(event_record["event_data"]), } insert_sql = f""" INSERT INTO {self._events_table} (id, app_name, user_id, session_id, invocation_id, timestamp, event_data) @@ -631,32 +616,32 @@ def _append_event(self, event_record: StoredEvent) -> None: def _delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._events_table} WHERE timestamp < @before" params: dict[str, Any] = {"before": before} + types: dict[str, Any] = {"before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(sql, params) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + types["app_name"] = SPANNER_PARAM_TYPES.STRING + return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) def _delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._session_table} WHERE update_time < @updated_before" params: dict[str, Any] = {"updated_before": updated_before} + types: dict[str, Any] = {"updated_before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(sql, params) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + types["app_name"] = SPANNER_PARAM_TYPES.STRING + return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) def _delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: sql = f"DELETE FROM {self._user_state_table} WHERE update_time < @updated_before" params: dict[str, Any] = {"updated_before": updated_before} + types: dict[str, Any] = {"updated_before": SPANNER_PARAM_TYPES.TIMESTAMP} if app_name is not None: sql += " AND app_name = @app_name" params["app_name"] = app_name - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(sql, params) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + types["app_name"] = SPANNER_PARAM_TYPES.STRING + return int(cast("Any", self._database()).run_in_transaction(_SpannerUpdateJob(sql, params, types))) def _get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = @app_name LIMIT 1" @@ -686,9 +671,7 @@ def _upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: INSERT OR UPDATE {self._app_state_table} (app_name, state, update_time) VALUES (@app_name, @state, PENDING_COMMIT_TIMESTAMP()) """ - self._run_write([ - (sql, {"app_name": app_name, "state": _to_spanner_json_payload(state)}, self._app_state_param_types()) - ]) + self._run_write([(sql, {"app_name": app_name, "state": to_json(state)}, self._app_state_param_types())]) def _upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: sql = f""" @@ -696,11 +679,7 @@ def _upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any] VALUES (@app_name, @user_id, @state, PENDING_COMMIT_TIMESTAMP()) """ self._run_write([ - ( - sql, - {"app_name": app_name, "user_id": user_id, "state": _to_spanner_json_payload(state)}, - self._user_state_param_types(), - ) + (sql, {"app_name": app_name, "user_id": user_id, "state": to_json(state)}, self._user_state_param_types()) ]) def _get_metadata(self, key: str) -> "str | None": @@ -900,14 +879,10 @@ def _run_read( return list(result_set) def _run_write(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: - with self._config.provide_session(transaction=True) as driver: - for sql, params, types in statements: - driver.execute(sql, _prepare_spanner_write_params(params, types)) + self._database().run_in_transaction(_SpannerMemoryWriteJob(statements)) # type: ignore[no-untyped-call] def _execute_update(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> int: - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(sql, _prepare_spanner_write_params(params, types)) - return int(getattr(result, "rows_affected", getattr(result, "rowcount", 0))) + return int(self._database().run_in_transaction(_SpannerMemoryUpdateJob(sql, params, types))) # type: ignore[no-untyped-call] def _memory_param_types(self, include_owner: bool) -> "dict[str, Any]": types: dict[str, Any] = { @@ -932,11 +907,8 @@ def _decode_json(self, raw: Any) -> Any: if raw is None: return None if isinstance(raw, str): - try: - return from_json(raw) - except Exception: - return raw - return _unwrap_spanner_json_object(raw) + return from_json(raw) + return raw def _create_tables(self) -> None: if not self._enabled: @@ -1169,53 +1141,6 @@ def _rows_to_records(self, rows: "list[Any]") -> "list[StoredMemory]": ] -def _prepare_spanner_write_params(params: "dict[str, Any]", types: "dict[str, Any] | None") -> "dict[str, Any]": - """Prepare ADK write parameters for Spanner driver execution.""" - if not types: - return params - json_type = _json_param_type() - changed = False - prepared: dict[str, Any] = {} - for key, value in params.items(): - param_type = types.get(key) - if param_type == json_type: - if value is None: - prepared[key] = TypedParameter(None, dict) - changed = True - elif isinstance(value, (str, bytes)): - prepared[key] = _to_spanner_json_payload(value) - changed = True - else: - prepared[key] = value - elif value is None and param_type is not None: - if param_type == SPANNER_PARAM_TYPES.STRING: - prepared[key] = TypedParameter(None, str) - changed = True - elif param_type == SPANNER_PARAM_TYPES.TIMESTAMP: - prepared[key] = TypedParameter(None, datetime) - changed = True - elif param_type == SPANNER_PARAM_TYPES.INT64: - prepared[key] = TypedParameter(None, int) - changed = True - else: - prepared[key] = value - else: - prepared[key] = value - return prepared if changed else params - - -def _to_spanner_json_payload(value: Any) -> Any: - """Prepare a value for Spanner JSON column parameter binding.""" - if value is None: - return None - if isinstance(value, (str, bytes)): - try: - return from_json(value) - except Exception: - return value - return value - - def _json_param_type() -> Any: try: return SPANNER_PARAM_TYPES.JSON @@ -1284,6 +1209,64 @@ def _spanner_drop_statement_table(statement: str, existing_tables: "set[str]") - return None +class _SpannerWriteJob: + __slots__ = ("_statements",) + + def __init__(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: + self._statements = statements + + def __call__(self, transaction: "Transaction") -> None: + if len(self._statements) > 1: + status, _row_counts = transaction.batch_update(self._statements) # type: ignore[no-untyped-call] + if status.code != 0: + msg = f"Spanner batch update failed (code {status.code}): {status.message}" + raise OperationalError(msg) + return + for sql, params, types in self._statements: + transaction.execute_update(sql, params=params, param_types=types) # type: ignore[no-untyped-call] + + +class _SpannerMemoryWriteJob: + __slots__ = ("_statements",) + + def __init__(self, statements: "list[tuple[str, dict[str, Any], dict[str, Any]]]") -> None: + self._statements = statements + + def __call__(self, transaction: "Transaction") -> None: + if len(self._statements) > 1: + status, _row_counts = transaction.batch_update(self._statements) # type: ignore[no-untyped-call] + if status.code != 0: + msg = f"Spanner batch update failed (code {status.code}): {status.message}" + raise OperationalError(msg) + return + for sql, params, types in self._statements: + transaction.execute_update(sql, params=params, param_types=types) # type: ignore[no-untyped-call] + + +class _SpannerUpdateJob: + __slots__ = ("_params", "_sql", "_types") + + def __init__(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: + self._sql = sql + self._params = params + self._types = types + + def __call__(self, transaction: "Transaction") -> int: + return int(transaction.execute_update(self._sql, params=self._params, param_types=self._types)) # type: ignore[no-untyped-call] + + +class _SpannerMemoryUpdateJob: + __slots__ = ("_params", "_sql", "_types") + + def __init__(self, sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: + self._sql = sql + self._params = params + self._types = types + + def __call__(self, transaction: "Transaction") -> int: + return int(transaction.execute_update(self._sql, params=self._params, param_types=self._types)) # type: ignore[no-untyped-call] + + class _SpannerReadProtocol(Protocol): def execute_sql( self, sql: str, params: "dict[str, Any] | None" = None, param_types: "dict[str, Any] | None" = None diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index ae2ac9710..aa475eafd 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -5,7 +5,6 @@ from typing_extensions import NotRequired -from sqlspec.adapters.spanner import _typing as spanner_typing from sqlspec.adapters.spanner._typing import SpannerConnection from sqlspec.adapters.spanner._typing import SpannerTransactionType as TransactionType from sqlspec.adapters.spanner.core import apply_driver_features, default_statement_config @@ -34,7 +33,6 @@ from sqlspec.adapters.spanner._typing import SpannerDirectedReadOptions as DirectedReadOptions from sqlspec.adapters.spanner._typing import SpannerEncryptionConfig as EncryptionConfig from sqlspec.adapters.spanner._typing import SpannerExecuteSqlRequest as ExecuteSqlRequest - from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool from sqlspec.adapters.spanner._typing import SpannerRequestOptions as RequestOptions from sqlspec.adapters.spanner._typing import SpannerRetry as Retry from sqlspec.config import ExtensionConfigs @@ -134,7 +132,7 @@ class SpannerConnectionParams(TypedDict): class SpannerPoolParams(SpannerConnectionParams): """Session pool configuration.""" - pool_type: "NotRequired[type[AbstractSessionPool] | str | None]" + pool_type: "NotRequired[type[AbstractSessionPool]]" size: "NotRequired[int]" target_size: "NotRequired[int]" max_sessions: "NotRequired[int]" @@ -142,9 +140,7 @@ class SpannerPoolParams(SpannerConnectionParams): session_labels: "NotRequired[dict[str, str]]" labels: "NotRequired[dict[str, str]]" ping_interval: "NotRequired[int]" - ping_timeout: "NotRequired[float]" max_age_minutes: "NotRequired[int]" - enable_multiplexed_sessions: "NotRequired[bool]" class SpannerDriverFeatures(TypedDict): @@ -178,7 +174,6 @@ class SpannerDriverFeatures(TypedDict): retry: "NotRequired[Retry | None]" timeout: "NotRequired[float | None]" request_options: "NotRequired[RequestOptions | dict[str, Any] | None]" - query_options: "NotRequired[ExecuteSqlRequest.QueryOptions | dict[str, Any] | None]" directed_read_options: "NotRequired[DirectedReadOptions | None]" session_labels: "NotRequired[dict[str, str]]" enable_events: "NotRequired[bool]" @@ -229,9 +224,10 @@ def __enter__(self) -> SpannerConnection: raise else: return self._connection - self._session = cast("Any", database).snapshot(multi_use=True) - self._connection = cast("SpannerConnection", self._session.__enter__()) - return self._connection + else: + self._session = cast("Any", database).snapshot(multi_use=True) + self._connection = cast("SpannerConnection", self._session.__enter__()) + return self._connection def __exit__( self, exc_type: "type[BaseException] | None", exc_val: "BaseException | None", exc_tb: "TracebackType | None" @@ -239,15 +235,25 @@ def __exit__( if self._transaction and self._connection: txn = cast("Any", self._connection) try: - rolled_back = bool(getattr(txn, "rolled_back", False)) - committed = getattr(txn, "committed", None) - txn_id = getattr(txn, "_transaction_id", None) - mutations = cast("list[Any] | None", getattr(txn, "_mutations", None)) if exc_type is None: - if not rolled_back and committed is None and (txn_id is not None or bool(mutations)): + try: + txn_id = txn._transaction_id + except AttributeError: + txn_id = None + mutations = cast("list[Any] | None", getattr(txn, "_mutations", None)) + try: + committed = txn.committed + except AttributeError: + committed = None + if committed is None and (txn_id is not None or bool(mutations)): txn.commit() - elif not rolled_back and committed is None and txn_id is not None: - txn.rollback() + else: + try: + rollback_txn_id = txn._transaction_id + except AttributeError: + rollback_txn_id = None + if rollback_txn_id is not None: + txn.rollback() finally: if self._session: cast("Any", self._config.get_database()).sessions_manager.put_session(self._session) @@ -318,13 +324,10 @@ def __init__( ): self.connection_config["session_labels"] = legacy_session_labels + from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool + self.connection_config.setdefault("size", self.connection_config.pop("max_sessions", 10)) - enable_multiplexed = self.connection_config.get("enable_multiplexed_sessions", True) - if enable_multiplexed and "pool_type" not in self.connection_config: - self.connection_config["pool_type"] = None - elif not enable_multiplexed and "pool_type" not in self.connection_config: - self.connection_config["pool_type"] = spanner_typing.SpannerPingingPool - self.connection_config.setdefault("ping_interval", 1800) + self.connection_config.setdefault("pool_type", FixedSizePool) statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, raw_driver_features) @@ -345,9 +348,11 @@ def __init__( self._database: Database | None = None def _get_client(self) -> "Client": + from sqlspec.adapters.spanner._typing import SpannerClient as Client + if self._client is None: client_kwargs = self._connection_kwargs_for(_CLIENT_CONFIG_FIELDS) - self._client = spanner_typing.SpannerClient(**client_kwargs) + self._client = Client(**client_kwargs) return self._client def get_database(self) -> "Database": @@ -357,11 +362,7 @@ def get_database(self) -> "Database": msg = "instance_id and database_id are required." raise ImproperConfigurationError(msg) - pool_type = self.connection_config.get("pool_type") - enable_multiplexed = self.connection_config.get("enable_multiplexed_sessions", True) - is_multiplexed = enable_multiplexed and (pool_type is None or pool_type == "multiplexed") - - if not is_multiplexed and self.connection_instance is None: + if self.connection_instance is None: self.connection_instance = self.provide_pool() if self._database is None: @@ -371,8 +372,7 @@ def get_database(self) -> "Database": if instance_labels is not None: instance_kwargs["labels"] = instance_labels database_kwargs = self._connection_kwargs_for(_DATABASE_CONFIG_FIELDS) - if self.connection_instance is not None: - database_kwargs["pool"] = self.connection_instance + database_kwargs["pool"] = self.connection_instance self._database = client.instance(instance_id, **instance_kwargs).database( # type: ignore[no-untyped-call] database_id, **database_kwargs ) @@ -391,27 +391,25 @@ def create_connection(self) -> SpannerConnection: return cast("SpannerConnection", self.get_database().snapshot(multi_use=True)) # type: ignore[no-untyped-call] def _create_pool(self) -> "AbstractSessionPool": + from sqlspec.adapters.spanner._typing import SpannerBurstyPool as BurstyPool + from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool + from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool + instance_id = self.connection_config.get("instance_id") database_id = self.connection_config.get("database_id") if not instance_id or not database_id: msg = "instance_id and database_id are required." raise ImproperConfigurationError(msg) - raw_pool_type = self.connection_config.get("pool_type") - pool_type: type[AbstractSessionPool | PingingPool] - if raw_pool_type is None or raw_pool_type == "multiplexed": - pool_type = spanner_typing.SpannerPingingPool - else: - pool_type = cast("type[AbstractSessionPool]", raw_pool_type) + pool_type = cast("type[AbstractSessionPool]", self.connection_config.get("pool_type", FixedSizePool)) labels = self.connection_config.get("session_labels", self.connection_config.get("labels")) pool_kwargs: dict[str, Any] = self._pool_base_kwargs(labels=cast("dict[str, str] | None", labels)) - if issubclass(pool_type, spanner_typing.SpannerPingingPool): - self.connection_config.setdefault("ping_interval", 1800) + if issubclass(pool_type, PingingPool): pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "ping_interval"})) - elif issubclass(pool_type, spanner_typing.SpannerFixedSizePool): + elif issubclass(pool_type, FixedSizePool): pool_kwargs.update(self._connection_kwargs_for({"size", "default_timeout", "max_age_minutes"})) - elif issubclass(pool_type, spanner_typing.SpannerBurstyPool): + elif issubclass(pool_type, BurstyPool): target_size = self.connection_config.get("target_size", self.connection_config.get("size")) if target_size is not None: pool_kwargs["target_size"] = target_size @@ -482,7 +480,6 @@ def provide_session( transaction: "bool" = _DEFAULT_SESSION_TRANSACTION, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, - query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -500,7 +497,6 @@ def provide_session( Snapshot (False). request_options: Session-scoped RequestOptions for Spanner statements. directed_read_options: Session-scoped DirectedReadOptions for reads. - query_options: Session-scoped QueryOptions for Spanner statements. retry: Session-scoped retry policy for Spanner statement calls. timeout: Session-scoped timeout for Spanner statement calls. **kwargs: Additional keyword arguments. @@ -518,7 +514,6 @@ def provide_session( driver_features=self._session_driver_features( request_options=request_options, directed_read_options=directed_read_options, - query_options=query_options, retry=retry, timeout=timeout, ), @@ -531,7 +526,6 @@ def provide_write_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, - query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -543,7 +537,6 @@ def provide_write_session( transaction=True, request_options=request_options, directed_read_options=directed_read_options, - query_options=query_options, retry=retry, timeout=timeout, **kwargs, @@ -555,7 +548,6 @@ def provide_read_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, - query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -571,7 +563,6 @@ def provide_read_session( transaction=False, request_options=request_options, directed_read_options=directed_read_options, - query_options=query_options, retry=retry, timeout=timeout, **kwargs, @@ -582,25 +573,16 @@ def _session_driver_features( *, request_options: "RequestOptions | dict[str, Any] | None", directed_read_options: "DirectedReadOptions | None", - query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None", timeout: "float | None", ) -> "dict[str, Any]": - if ( - request_options is None - and directed_read_options is None - and query_options is None - and retry is None - and timeout is None - ): + if request_options is None and directed_read_options is None and retry is None and timeout is None: return self.driver_features driver_features = dict(self.driver_features) if request_options is not None: driver_features["request_options"] = request_options if directed_read_options is not None: driver_features["directed_read_options"] = directed_read_options - if query_options is not None: - driver_features["query_options"] = query_options if retry is not None: driver_features["retry"] = retry if timeout is not None: diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index e12c9f631..040b3ae3e 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -333,15 +333,10 @@ def _convert_json_row_value(value: Any, *, json_deserializer: "Callable[[str], A if isinstance(value, JsonObject): if json_deserializer is from_json: return _unwrap_spanner_json_object(value) - if getattr(value, "_is_null", False): - return None try: - serialized = cast("Any", value).serialize() - if serialized is None: - return None - return json_deserializer(serialized) + return json_deserializer(cast("Any", value).serialize()) except (TypeError, ValueError): - return _unwrap_spanner_json_object(value) + return value elif isinstance(value, str): try: return json_deserializer(value) diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 2962a87d3..322426006 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -28,7 +28,6 @@ ) from sqlspec.adapters.spanner.data_dictionary import SpannerDataDictionary from sqlspec.core import StatementConfig, register_driver_profile -from sqlspec.core.statement import SQL from sqlspec.driver import ( BaseSyncExceptionHandler, ExecutionResult, @@ -46,11 +45,11 @@ from sqlspec.adapters.spanner._typing import SpannerConnection from sqlspec.adapters.spanner._typing import SpannerDirectedReadOptions as DirectedReadOptions - from sqlspec.adapters.spanner._typing import SpannerExecuteSqlRequest as ExecuteSqlRequest from sqlspec.adapters.spanner._typing import SpannerRequestOptions as RequestOptions from sqlspec.adapters.spanner._typing import SpannerRetry as Retry from sqlspec.builder import QueryBuilder from sqlspec.core import ArrowResult, SQLResult, Statement, StatementFilter + from sqlspec.core.statement import SQL from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import SchemaT, StatementParameters @@ -94,7 +93,7 @@ class SpannerSyncDriver(SyncDriverAdapterBase): """Synchronous Spanner driver operating on Snapshot or Transaction contexts.""" dialect: "DialectType" = "spanner" - __slots__ = ("_config", "_data_dictionary", "_pending_execute_options", "_row_plan_cache", "_row_plan_deserializer") + __slots__ = ("_data_dictionary", "_pending_execute_options", "_row_plan_cache", "_row_plan_deserializer") def __init__( self, @@ -107,7 +106,6 @@ def __init__( statement_config = default_statement_config super().__init__(connection=connection, statement_config=statement_config, driver_features=features) - self._config: Any = None self._data_dictionary: SpannerDataDictionary | None = None self._pending_execute_options: _PerCallExecuteOptions | None = None self._row_plan_cache: dict[int, tuple[Any, list[str], tuple[tuple[int, Any], ...] | None]] = {} @@ -178,7 +176,7 @@ def dispatch_execute_many(self, cursor: "SpannerConnection", statement: "SQL") - _coerce = self._coerce_params _infer = self._infer_param_types - execute_kwargs = self._execute_kwargs(for_batch=True) + execute_kwargs = self._execute_kwargs() param_types_cache: dict[tuple[tuple[str, type[Any], Any], ...], dict[str, Any]] = {} empty_param_types: dict[str, Any] = {} batch_args: list[tuple[str, dict[str, Any] | None, dict[str, Any]]] = [] @@ -239,35 +237,15 @@ def begin(self) -> None: return None def commit(self) -> None: - writer = cast("_SpannerWriteProtocol", self.connection) - if not ( - isinstance(self.connection, SpannerTransaction) - or supports_write(self.connection) - or callable(getattr(writer, "commit", None)) - ): - return - if getattr(writer, "rolled_back", False) or getattr(writer, "committed", None) is not None: - return - if ( - getattr(writer, "_transaction_id", None) is not None - or bool(getattr(writer, "_mutations", None)) - or (not hasattr(writer, "_transaction_id") and not hasattr(writer, "_mutations")) - ) and callable(getattr(writer, "commit", None)): + if isinstance(self.connection, SpannerTransaction): + writer = cast("_SpannerWriteProtocol", self.connection) + if writer.committed is not None: + return writer.commit() def rollback(self) -> None: - writer = cast("_SpannerWriteProtocol", self.connection) - if not ( - isinstance(self.connection, SpannerTransaction) - or supports_write(self.connection) - or callable(getattr(writer, "rollback", None)) - ): - return - if getattr(writer, "rolled_back", False) or getattr(writer, "committed", None) is not None: - return - if (getattr(writer, "_transaction_id", None) is not None or not hasattr(writer, "_transaction_id")) and ( - callable(getattr(writer, "rollback", None)) - ): + if isinstance(self.connection, SpannerTransaction): + writer = cast("_SpannerWriteProtocol", self.connection) writer.rollback() def create_savepoint(self, name: str) -> None: @@ -303,79 +281,6 @@ def with_cursor(self, connection: "SpannerConnection") -> "SpannerSyncCursor": def handle_database_exceptions(self) -> "SpannerExceptionHandler": return SpannerExceptionHandler() - def _get_database(self) -> Any: - if self._config is not None: - return self._config.get_database() - session = getattr(self.connection, "_session", None) - if session is not None: - database = getattr(session, "_database", None) - if database is not None: - return database - database = getattr(self.connection, "_database", None) - if database is not None: - return database - return None - - def _execute_partitioned_dml( - self, - statement: "SQL | str", - *parameters: Any, - query_options: Any = None, - request_options: Any = None, - exclude_txn_from_change_streams: bool = False, - **kwargs: Any, - ) -> int: - """Execute a Partitioned DML statement across database partitions. - - Args: - statement: The SQL string or SQL object to execute. - *parameters: Positional parameters or parameter mapping. - query_options: Optional Spanner QueryOptions. - request_options: Optional Spanner RequestOptions. - exclude_txn_from_change_streams: Whether to exclude the transaction from change streams. - **kwargs: Additional keyword arguments or parameters. - - Returns: - The number of affected rows. - """ - database = self._get_database() - if database is None: - msg = "Could not resolve Spanner database for partitioned DML execution." - raise SQLConversionError(msg) - - if isinstance(statement, SQL): - sql_statement = statement - else: - sql_statement = self.prepare_statement( - statement, parameters, statement_config=self.statement_config, kwargs=kwargs or None - ) - - sql, raw_params = self._compiled_sql(sql_statement, self.statement_config) - params = raw_params if isinstance(raw_params, dict) else None - coerced_params = self._coerce_params(params) - param_types = self._infer_param_types(params) - - effective_request_options = request_options or self.driver_features.get("request_options") - effective_query_options = query_options or self.driver_features.get("query_options") - - call_kwargs: dict[str, Any] = { - "params": coerced_params, - "param_types": param_types, - "exclude_txn_from_change_streams": exclude_txn_from_change_streams, - } - if effective_query_options is not None: - call_kwargs["query_options"] = effective_query_options - if effective_request_options is not None: - call_kwargs["request_options"] = effective_request_options - - result: Any = 0 - exc_handler = self.handle_database_exceptions() - with exc_handler: - result = database.execute_partitioned_dml(sql, **call_kwargs) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - return int(result) - def execute( self, statement: "SQL | Statement | QueryBuilder", @@ -534,18 +439,24 @@ def load_from_arrow( arrow_table = self._coerce_arrow_table(source) if overwrite: - self._execute_partitioned_dml(f"DELETE FROM {table} WHERE TRUE") + delete_sql = f"DELETE FROM {table} WHERE TRUE" + if isinstance(self.connection, SpannerTransaction): + writer = cast("_SpannerWriteProtocol", self.connection) + writer.execute_update(delete_sql) + else: + msg = "Delete requires a Transaction context." + raise SQLConversionError(msg) columns, records = self._arrow_table_to_rows(arrow_table) if records: + conn = self.connection + if not isinstance(conn, SpannerTransaction): + msg = "Arrow import requires a Transaction context." + raise SQLConversionError(msg) chunks = self._chunk_mutation_rows(columns, records) - if self.driver_features.get("enable_batch_write_api"): + if self.driver_features.get("enable_batch_write_api") and not overwrite: self._batch_write_mutations(table, columns, chunks) else: - conn = self.connection - if not isinstance(conn, SpannerTransaction): - msg = "Arrow import requires a Transaction context." - raise SQLConversionError(msg) writer = cast("_SpannerWriteProtocol", conn) for chunk in chunks: writer.insert_or_update(table, columns, chunk) @@ -601,57 +512,36 @@ def resolve_rowcount(self, cursor: "SpannerConnection") -> int: """ return 0 - def _execute_kwargs(self, *, for_read: bool = False, for_batch: bool = False) -> dict[str, Any]: + def _execute_kwargs(self, *, for_read: bool = False) -> dict[str, Any]: kwargs: dict[str, Any] = { key: self.driver_features[key] for key in ("retry", "timeout") if key in self.driver_features } request_options = self.driver_features.get("request_options") if request_options is not None: kwargs["request_options"] = request_options - if not for_batch: - query_options = self.driver_features.get("query_options") - if query_options is not None: - kwargs["query_options"] = query_options - if for_read and not for_batch: - directed_read_options = self.driver_features.get("directed_read_options") - if directed_read_options is not None: - kwargs["directed_read_options"] = directed_read_options + directed_read_options = self.driver_features.get("directed_read_options") + if for_read and directed_read_options is not None: + kwargs["directed_read_options"] = directed_read_options pending = self._pending_execute_options if pending is not None: if pending.request_options is not None: kwargs["request_options"] = pending.request_options - if not for_batch and pending.query_options is not None: - kwargs["query_options"] = pending.query_options if pending.retry is not None: kwargs["retry"] = pending.retry if pending.timeout is not None: kwargs["timeout"] = pending.timeout - if for_read and not for_batch and pending.directed_read_options is not None: + if for_read and pending.directed_read_options is not None: kwargs["directed_read_options"] = pending.directed_read_options - if not for_read and pending.last_statement: - kwargs["last_statement"] = True return kwargs def _pop_execute_options(self, kwargs: dict[str, Any]) -> "_PerCallExecuteOptions | None": - if not any( - key in kwargs - for key in ( - "request_options", - "query_options", - "directed_read_options", - "retry", - "timeout", - "last_statement", - ) - ): + if not any(key in kwargs for key in ("request_options", "directed_read_options", "retry", "timeout")): return None return _PerCallExecuteOptions( request_options=kwargs.pop("request_options", None), - query_options=kwargs.pop("query_options", None), directed_read_options=kwargs.pop("directed_read_options", None), retry=kwargs.pop("retry", None), timeout=kwargs.pop("timeout", None), - last_statement=bool(kwargs.pop("last_statement", False)), ) def _chunk_mutation_rows(self, columns: "list[str]", records: "list[tuple[Any, ...]]") -> "list[list[list[Any]]]": @@ -679,9 +569,10 @@ def _chunk_mutation_rows(self, columns: "list[str]", records: "list[tuple[Any, . def _batch_write_mutations(self, table: str, columns: "list[str]", chunks: "list[list[list[Any]]]") -> None: """High-throughput ingest via the Spanner Batch Write API (one mutation group per chunk).""" - database = self._get_database() + session = cast("object", getattr(self.connection, "_session", None)) + database = cast("Any", getattr(session, "_database", None)) if session is not None else None if database is None: - msg = "Spanner Batch Write API requires a database-backed session or config." + msg = "Spanner Batch Write API requires a database-backed session." raise SQLConversionError(msg) with database.mutation_groups() as mutation_groups: for chunk in chunks: @@ -756,24 +647,20 @@ def rollback(self) -> None: ... class _PerCallExecuteOptions: """Per-call Spanner execution options captured for a single dispatch.""" - __slots__ = ("directed_read_options", "last_statement", "query_options", "request_options", "retry", "timeout") + __slots__ = ("directed_read_options", "request_options", "retry", "timeout") def __init__( self, *, request_options: "RequestOptions | dict[str, Any] | None" = None, - query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, - last_statement: bool = False, ) -> None: self.request_options = request_options - self.query_options = query_options self.directed_read_options = directed_read_options self.retry = retry self.timeout = timeout - self.last_statement = last_statement class _SpannerSelectStreamSource: diff --git a/sqlspec/adapters/spanner/litestar/store.py b/sqlspec/adapters/spanner/litestar/store.py index ab6243a99..e81ecb153 100644 --- a/sqlspec/adapters/spanner/litestar/store.py +++ b/sqlspec/adapters/spanner/litestar/store.py @@ -5,15 +5,26 @@ from typing_extensions import NotRequired -from sqlspec.adapters.spanner.type_converter import spanner_to_bytes +from sqlspec.adapters.spanner._typing import spanner_param_types as param_types +from sqlspec.adapters.spanner.type_converter import bytes_to_spanner, spanner_to_bytes from sqlspec.config import LitestarConfig -from sqlspec.core import TypedParameter from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ if TYPE_CHECKING: + from collections.abc import Callable + from typing import Protocol + + from sqlspec.adapters.spanner._typing import SpannerTransaction as Transaction from sqlspec.adapters.spanner.config import SpannerSyncConfig + class _DatabaseProtocol(Protocol): + def run_in_transaction(self, func: "Callable[[Transaction], Any]") -> Any: ... + + def update_ddl(self, ddl_statements: "list[str]") -> Any: ... + + def list_tables(self) -> Any: ... + __all__ = ("SpannerLitestarConfig", "SpannerSyncStore") @@ -79,6 +90,9 @@ async def create_table(self) -> None: await async_(self._create_table)() await self.reconcile_schema(assume_existing=True) + def _database(self) -> "_DatabaseProtocol": + return cast("_DatabaseProtocol", self._config.get_database()) + def _datetime_to_timestamp(self, dt: "datetime | None") -> "datetime | None": if dt is None: return None @@ -96,14 +110,23 @@ def _timestamp_to_datetime(self, ts: "datetime | None") -> "datetime | None": def _build_params( self, key: str, expires_at: "datetime | None" = None, data: "bytes | None" = None ) -> "dict[str, Any]": - ts = self._datetime_to_timestamp(expires_at) - params: dict[str, Any] = { + return { "session_id": key, - "expires_at": ts if ts is not None else TypedParameter(None, datetime), + "data": bytes_to_spanner(data), + "expires_at": self._datetime_to_timestamp(expires_at), } - if data is not None: - params["data"] = data - return params + + def _get_param_types( + self, session_id: bool = True, expires_at: bool = False, data: bool = False + ) -> "dict[str, Any]": + types: dict[str, Any] = {} + if session_id: + types["session_id"] = param_types.STRING + if expires_at: + types["expires_at"] = param_types.TIMESTAMP + if data: + types["data"] = param_types.BYTES + return types def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": sql = f""" @@ -140,8 +163,8 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | if self._shard_count > 1: update_sql = f"{update_sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = self._build_params(key, new_expires) - with self._config.provide_session(transaction=True) as driver: - driver.execute(update_sql, params) + types = self._get_param_types(expires_at=True) + self._database().run_in_transaction(_SpannerExecuteUpdateJob(update_sql, params, types)) return spanner_to_bytes(data) @@ -149,6 +172,7 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) params = self._build_params(key, expires_at, data) + types = self._get_param_types(session_id=True, expires_at=True, data=True) update_sql = f""" UPDATE {self._table_name} @@ -163,26 +187,19 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No INSERT {self._table_name} (session_id, data, expires_at, created_at, updated_at) VALUES (@session_id, @data, @expires_at, PENDING_COMMIT_TIMESTAMP(), PENDING_COMMIT_TIMESTAMP()) """ - - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(update_sql, params) - rows_affected = getattr(result, "rows_affected", None) - has_rows = rows_affected > 0 if isinstance(rows_affected, int) else bool(getattr(result, "rowcount", None)) - if not has_rows: - driver.execute(insert_sql, params) + self._database().run_in_transaction(_SpannerUpsertJob(update_sql, insert_sql, params, types)) def _delete(self, key: str) -> None: sql = f"DELETE FROM {self._table_name} WHERE session_id = @session_id" if self._shard_count > 1: sql = f"{sql} AND shard_id = MOD(FARM_FINGERPRINT(@session_id), {self._shard_count})" params = {"session_id": key} - with self._config.provide_session(transaction=True) as driver: - driver.execute(sql, params) + types = self._get_param_types(session_id=True) + self._database().run_in_transaction(_SpannerExecuteUpdateJob(sql, params, types)) def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name} WHERE TRUE" - with self._config.provide_session(transaction=True) as driver: - driver.execute(sql) + self._database().run_in_transaction(_SpannerExecuteUpdateJob(sql)) def _exists(self, key: str) -> bool: sql = f""" @@ -217,12 +234,8 @@ def _delete_expired(self) -> int: DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= CURRENT_TIMESTAMP() """ - with self._config.provide_session(transaction=True) as driver: - result = driver.execute(sql) - rows_affected = getattr(result, "rows_affected", None) - if isinstance(rows_affected, int): - return rows_affected - return int(getattr(result, "rowcount", 0)) + result = self._database().run_in_transaction(_SpannerExecuteUpdateCountJob(sql)) + return cast("int", result) def _create_table(self) -> None: database = self._config.get_database() @@ -262,3 +275,43 @@ def _index_ddl(self) -> str: def _drop_table_sql(self) -> "list[str]": return [f"DROP INDEX idx_{self._table_name}_expires_at", f"DROP TABLE {self._table_name}"] + + +class _SpannerExecuteUpdateJob: + __slots__ = ("_params", "_sql", "_types") + + def __init__(self, sql: str, params: "dict[str, Any] | None" = None, types: "dict[str, Any] | None" = None) -> None: + self._sql = sql + self._params = params + self._types = types + + def __call__(self, transaction: "Transaction") -> None: + if self._params is None and self._types is None: + transaction.execute_update(self._sql) # type: ignore[no-untyped-call] + return + transaction.execute_update(self._sql, params=self._params or {}, param_types=self._types) # type: ignore[no-untyped-call] + + +class _SpannerUpsertJob: + __slots__ = ("_insert_sql", "_params", "_types", "_update_sql") + + def __init__(self, update_sql: str, insert_sql: str, params: "dict[str, Any]", types: "dict[str, Any]") -> None: + self._update_sql = update_sql + self._insert_sql = insert_sql + self._params = params + self._types = types + + def __call__(self, transaction: "Transaction") -> None: + row_ct = transaction.execute_update(self._update_sql, params=self._params, param_types=self._types) # type: ignore[no-untyped-call] + if row_ct == 0: + transaction.execute_update(self._insert_sql, params=self._params, param_types=self._types) # type: ignore[no-untyped-call] + + +class _SpannerExecuteUpdateCountJob: + __slots__ = ("_sql",) + + def __init__(self, sql: str) -> None: + self._sql = sql + + def __call__(self, transaction: "Transaction") -> int: + return int(transaction.execute_update(self._sql)) # type: ignore[no-untyped-call] diff --git a/sqlspec/adapters/spanner/type_converter.py b/sqlspec/adapters/spanner/type_converter.py index 46f00f0cd..11513b4be 100644 --- a/sqlspec/adapters/spanner/type_converter.py +++ b/sqlspec/adapters/spanner/type_converter.py @@ -17,7 +17,7 @@ """ import base64 -from datetime import date, datetime, timedelta, timezone +from datetime import date, datetime, timezone from decimal import Decimal from typing import TYPE_CHECKING, Any, cast from uuid import UUID @@ -33,8 +33,6 @@ from sqlspec.protocols import SpannerParamTypesProtocol __all__ = ( - "SPANNER_FLOAT32", - "SPANNER_VECTOR", "bytes_to_spanner", "coerce_params_for_spanner", "infer_spanner_param_types", @@ -44,9 +42,6 @@ "uuid_to_spanner", ) -SPANNER_FLOAT32: str = "FLOAT32" -SPANNER_VECTOR: str = "ARRAY" - _UUID_TYPES: "tuple[type[Any], ...]" = (UUID,) _uuid_utils_uuid = import_optional_attr("uuid_utils", "UUID") if _uuid_utils_uuid is not None: @@ -173,18 +168,10 @@ def coerce_params_for_spanner( coerced: dict[str, Any] = {} changed = False for key, value in params.items(): - declared_type = None if type(value) is TypedParameter: - declared_type = value.original_type value = value.value changed = True - if declared_type in ("ARRAY", "VECTOR", "vector") and isinstance(value, (list, tuple)): - coerced[key] = [float(x) for x in value] - changed = True - elif declared_type in ("FLOAT32", "float32") and value is not None: - coerced[key] = float(value) - changed = True - elif isinstance(value, _UUID_TYPES): + if isinstance(value, _UUID_TYPES): if enable_uuid_conversion: coerced[key] = str(value) changed = True @@ -215,7 +202,7 @@ def coerce_params_for_spanner( return coerced if changed else params -_NULL_PARAM_TYPE_NAMES: "dict[type[Any] | str, str]" = { +_NULL_PARAM_TYPE_NAMES: "dict[type[Any], str]" = { bool: "BOOL", int: "INT64", float: "FLOAT64", @@ -224,17 +211,8 @@ def coerce_params_for_spanner( datetime: "TIMESTAMP", date: "DATE", Decimal: "NUMERIC", - timedelta: "INTERVAL", UUID: "STRING", dict: "JSON", - "FLOAT32": "FLOAT32", - "float32": "FLOAT32", - "FLOAT64": "FLOAT64", - "float64": "FLOAT64", - "INTERVAL": "INTERVAL", - "interval": "INTERVAL", - "JSON": "JSON", - "json": "JSON", } @@ -251,24 +229,19 @@ def _infer_sequence_param_type(value: Any, param_types: Any, json_type: Any) -> """ if should_json_encode_sequence(value): return json_type - sequence = list(value) - if not sequence: + if not value: return None - first = sequence[0] + first = value[0] + if isinstance(first, bool): + return param_types.Array(param_types.BOOL) if isinstance(first, int): return param_types.Array(param_types.INT64) if isinstance(first, str): return param_types.Array(param_types.STRING) if isinstance(first, float): return param_types.Array(param_types.FLOAT64) - if isinstance(first, bool): - return param_types.Array(param_types.BOOL) if isinstance(first, Decimal): return param_types.Array(param_types.NUMERIC) - if isinstance(first, timedelta): - interval_type = getattr(param_types, "INTERVAL", None) - if interval_type is not None: - return param_types.Array(interval_type) return None @@ -291,25 +264,11 @@ def infer_spanner_param_types(params: "dict[str, Any] | None") -> "dict[str, Any for key, raw_value in params.items(): is_typed = type(raw_value) is TypedParameter value = raw_value.value if is_typed else raw_value - declared = raw_value.original_type if is_typed else None if value is None: null_type = _null_param_type(raw_value, param_types) if null_type is not None: types[key] = null_type continue - if declared in ("FLOAT32", "float32"): - types[key] = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) - continue - if declared in ("ARRAY", "VECTOR", "vector"): - float32_type = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) - types[key] = param_types.Array(float32_type) - continue - if declared in ("FLOAT64", "float64"): - types[key] = param_types.FLOAT64 - continue - if declared == "ARRAY": - types[key] = param_types.Array(param_types.FLOAT64) - continue if isinstance(value, bool): types[key] = param_types.BOOL elif isinstance(value, int): @@ -318,10 +277,6 @@ def infer_spanner_param_types(params: "dict[str, Any] | None") -> "dict[str, Any types[key] = param_types.FLOAT64 elif isinstance(value, Decimal): types[key] = param_types.NUMERIC - elif isinstance(value, timedelta): - interval_type = getattr(param_types, "INTERVAL", None) - if interval_type is not None: - types[key] = interval_type elif isinstance(value, _STRING_PARAM_TYPES): types[key] = param_types.STRING elif isinstance(value, bytes): @@ -357,16 +312,7 @@ def _null_param_type(raw_value: Any, param_types: "SpannerParamTypesProtocol") - declared = raw_value.original_type if type(raw_value) is TypedParameter else None if declared is None: return None - if declared in ("ARRAY", "VECTOR", "vector"): - float32_type = getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) - return param_types.Array(float32_type) - if declared == "ARRAY": - return param_types.Array(param_types.FLOAT64) resolver = _NULL_PARAM_TYPE_NAMES.get(declared) - if resolver == "FLOAT32": - return getattr(param_types, "FLOAT32", getattr(param_types, "FLOAT64", None)) - if resolver == "INTERVAL": - return getattr(param_types, "INTERVAL", None) if resolver == "JSON": return _json_param_type() return getattr(param_types, resolver) if resolver is not None else None diff --git a/sqlspec/core/parameters/_types.py b/sqlspec/core/parameters/_types.py index f56491a60..703cd9fb5 100644 --- a/sqlspec/core/parameters/_types.py +++ b/sqlspec/core/parameters/_types.py @@ -176,11 +176,9 @@ class TypedParameter: __slots__ = TYPED_PARAMETER_SLOTS - def __init__( - self, value: Any, original_type: "type | str | None" = None, semantic_name: "str | None" = None - ) -> None: + def __init__(self, value: Any, original_type: "type | None" = None, semantic_name: "str | None" = None) -> None: self.value = value - self.original_type: Any = original_type or type(value) + self.original_type = original_type or type(value) self.semantic_name = semantic_name self._hash: int | None = None @@ -201,8 +199,7 @@ def __eq__(self, other: object) -> bool: def __repr__(self) -> str: name_part = f", semantic_name='{self.semantic_name}'" if self.semantic_name else "" - type_name = getattr(self.original_type, "__name__", str(self.original_type)) - return f"TypedParameter({self.value!r}, original_type={type_name}{name_part})" + return f"TypedParameter({self.value!r}, original_type={self.original_type.__name__}{name_part})" def __reduce__(self) -> "tuple[Any, ...]": """Reconstruct via ``TypedParameter(value, original_type, semantic_name)``.""" diff --git a/sqlspec/protocols.py b/sqlspec/protocols.py index 9d85d8d5a..c9874ae44 100644 --- a/sqlspec/protocols.py +++ b/sqlspec/protocols.py @@ -282,10 +282,8 @@ class SpannerParamTypesProtocol(SupportsJsonTypeProtocol, Protocol): BOOL: Any INT64: Any - FLOAT32: Any FLOAT64: Any NUMERIC: Any - INTERVAL: Any STRING: Any BYTES: Any TIMESTAMP: Any diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index aa561ab7c..540f43c86 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -248,7 +248,6 @@ class SourceEquivalenceCase: "enable_events", "events_backend", "enable_batch_write_api", - "query_options", ), "sqlite": ( "enable_custom_adapters", diff --git a/tests/unit/adapters/test_spanner/test_batch_write_api.py b/tests/unit/adapters/test_spanner/test_batch_write_api.py index 198225c25..8a002e02a 100644 --- a/tests/unit/adapters/test_spanner/test_batch_write_api.py +++ b/tests/unit/adapters/test_spanner/test_batch_write_api.py @@ -53,15 +53,10 @@ def batch_write(self, request_options: Any = None, exclude_txn_from_change_strea class _FakeDatabase: def __init__(self) -> None: self.mutation_groups_obj = _FakeMutationGroups() - self.partitioned_dml_calls: list[str] = [] def mutation_groups(self) -> _FakeMutationGroups: return self.mutation_groups_obj - def execute_partitioned_dml(self, dml: str, **kwargs: Any) -> int: - self.partitioned_dml_calls.append(dml) - return 0 - class _FakeSession: def __init__(self, database: _FakeDatabase) -> None: @@ -115,13 +110,10 @@ def test_batch_write_splits_before_crossing_mutation_group_cell_cap(batch_write_ assert [len(chunk) for chunk in chunks] == [26_666, 1] -def test_batch_write_overwrite_uses_batch_write_after_partitioned_dml(batch_write_driver: SpannerSyncDriver) -> None: +def test_batch_write_overwrite_uses_transactional_mutations(batch_write_driver: SpannerSyncDriver) -> None: conn = cast("_FakeBatchTransaction", batch_write_driver.connection) batch_write_driver.load_from_arrow("users", pa.table({"id": [1]}), overwrite=True) - assert ( - conn.database.partitioned_dml_calls and "DELETE FROM users WHERE TRUE" in conn.database.partitioned_dml_calls[0] - ) - assert conn.insert_or_update_calls == [] - assert conn.database.mutation_groups_obj.batch_write_calls == 1 - assert conn.database.mutation_groups_obj.groups[0].calls == [("users", ["id"], [[1]])] + assert conn.execute_update_calls and "DELETE FROM users WHERE TRUE" in conn.execute_update_calls[0] + assert conn.insert_or_update_calls == [("users", ["id"], [[1]])] + assert conn.database.mutation_groups_obj.batch_write_calls == 0 diff --git a/tests/unit/adapters/test_spanner/test_config.py b/tests/unit/adapters/test_spanner/test_config.py index d8ae400c8..3fa465353 100644 --- a/tests/unit/adapters/test_spanner/test_config.py +++ b/tests/unit/adapters/test_spanner/test_config.py @@ -527,81 +527,6 @@ def __init__(self) -> None: assert db.session_obj.txn.rollback_calls == 0 -def test_transaction_context_and_driver_respect_rolled_back_state() -> None: - """Explicit driver.rollback() must prevent auto-commit or duplicate rollback on exit.""" - - class _Txn: - def __init__(self) -> None: - self._transaction_id: str | None = "txn-1" - self._mutations: list[object] = [object()] - self.committed: object | None = None - self.rolled_back = False - self.commit_calls = 0 - self.rollback_calls = 0 - - def __enter__(self): - return self - - def __exit__(self, *_: object) -> None: - return None - - def commit(self) -> None: - self.commit_calls += 1 - self.committed = object() - - def rollback(self) -> None: - if self.rolled_back: - msg = "Transaction already rolled back." - raise ValueError(msg) - self.rollback_calls += 1 - self.rolled_back = True - - class _Session: - def __init__(self) -> None: - self.txn = _Txn() - - def transaction(self) -> _Txn: - return self.txn - - class _SessionsManager: - def __init__(self, session: _Session) -> None: - self.session = session - self.returned = 0 - - def get_session(self, _transaction_type: object) -> _Session: - return self.session - - def put_session(self, _session: object) -> None: - self.returned += 1 - - class _DB: - def __init__(self) -> None: - self.session_obj = _Session() - self.sessions_manager = _SessionsManager(self.session_obj) - - db = _DB() - config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) - setattr(config, "get_database", lambda: db) - - with config.provide_session(transaction=True) as driver: - driver.rollback() - driver.rollback() - driver.commit() - - assert db.session_obj.txn.rollback_calls == 1 - assert db.session_obj.txn.commit_calls == 0 - - db_err = _DB() - setattr(config, "get_database", lambda: db_err) - with pytest.raises(RuntimeError, match="abort"), config.provide_session(transaction=True) as driver: - driver.rollback() - msg = "abort" - raise RuntimeError(msg) - - assert db_err.session_obj.txn.rollback_calls == 1 - assert db_err.session_obj.txn.commit_calls == 0 - - def test_provide_session_uses_batch_when_transaction_requested() -> None: """Driver should receive transaction connection when transaction=True.""" diff --git a/tests/unit/adapters/test_spanner/test_litestar_store.py b/tests/unit/adapters/test_spanner/test_litestar_store.py index ad2432f55..6022cc66c 100644 --- a/tests/unit/adapters/test_spanner/test_litestar_store.py +++ b/tests/unit/adapters/test_spanner/test_litestar_store.py @@ -1,72 +1,70 @@ -from datetime import datetime, timezone from typing import Any -from unittest.mock import MagicMock, call +from unittest.mock import MagicMock from sqlspec.adapters.spanner.litestar import SpannerSyncStore -from sqlspec.adapters.spanner.type_converter import bytes_to_spanner -def test_set_uses_session() -> None: - """Verify _set uses config.provide_session(transaction=True) for write operations.""" - driver = MagicMock() - driver.execute.return_value = MagicMock(rows_affected=1) - cm = _context_manager_yielding(driver) +def _mock_database() -> MagicMock: + """Create a mock database that captures run_in_transaction calls.""" + db = MagicMock() + db.run_in_transaction = MagicMock(side_effect=lambda func: func(MagicMock())) + return db + + +def test_set_uses_run_in_transaction() -> None: + """Verify _set uses database.run_in_transaction for write operations.""" + mock_db = _mock_database() config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.provide_session.return_value = cm + config.get_database.return_value = mock_db store = SpannerSyncStore(config) - store._set("s1", b"data", None) + store._set("s1", b"data", None) # pyright: ignore - config.provide_session.assert_called_once_with(transaction=True) + mock_db.run_in_transaction.assert_called_once() -def test_delete_uses_session() -> None: - """Verify _delete uses config.provide_session(transaction=True) for write operations.""" - driver = MagicMock() - cm = _context_manager_yielding(driver) +def test_delete_uses_run_in_transaction() -> None: + """Verify _delete uses database.run_in_transaction for write operations.""" + mock_db = _mock_database() config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.provide_session.return_value = cm + config.get_database.return_value = mock_db store = SpannerSyncStore(config) - store._delete("s1") + store._delete("s1") # pyright: ignore - config.provide_session.assert_called_once_with(transaction=True) + mock_db.run_in_transaction.assert_called_once() -def test_delete_all_uses_session() -> None: - """Verify _delete_all uses config.provide_session(transaction=True) for write operations.""" - driver = MagicMock() - cm = _context_manager_yielding(driver) +def test_delete_all_uses_run_in_transaction() -> None: + """Verify _delete_all uses database.run_in_transaction for write operations.""" + mock_db = _mock_database() config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.provide_session.return_value = cm + config.get_database.return_value = mock_db store = SpannerSyncStore(config) - store._delete_all() + store._delete_all() # pyright: ignore - config.provide_session.assert_called_once_with(transaction=True) + mock_db.run_in_transaction.assert_called_once() -def test_delete_expired_uses_session() -> None: - """Verify _delete_expired uses config.provide_session(transaction=True) for write operations.""" - driver = MagicMock() - driver.execute.return_value = MagicMock(rows_affected=3) - cm = _context_manager_yielding(driver) +def test_delete_expired_uses_run_in_transaction() -> None: + """Verify _delete_expired uses database.run_in_transaction for write operations.""" + mock_db = _mock_database() config = MagicMock() config.extension_config = {"litestar": {"session_table": "sess"}} - config.provide_session.return_value = cm + config.get_database.return_value = mock_db store = SpannerSyncStore(config) - result = store._delete_expired() + store._delete_expired() # pyright: ignore - config.provide_session.assert_called_once_with(transaction=True) - assert result == 3 + mock_db.run_in_transaction.assert_called_once() def _context_manager_yielding(value: Any) -> Any: @@ -91,32 +89,12 @@ def test_get_uses_snapshot_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - result = store._get("s1") + result = store._get("s1") # pyright: ignore config.provide_session.assert_called_once_with() assert result is None -def test_get_renewal_uses_transaction_session() -> None: - """Verify _get token renewal uses provide_session(transaction=True) for the UPDATE.""" - driver = MagicMock() - driver.select_one_or_none.return_value = { - "data": bytes_to_spanner(b"val"), - "expires_at": datetime(2030, 1, 1, tzinfo=timezone.utc), - } - cm = _context_manager_yielding(driver) - - config = MagicMock() - config.extension_config = {"litestar": {"session_table": "sess"}} - config.provide_session.return_value = cm - - store = SpannerSyncStore(config) - result = store._get("s1", renew_for=60) - - assert result == b"val" - assert config.provide_session.call_args_list == [call(), call(transaction=True)] - - def test_exists_uses_snapshot_session() -> None: """Verify _exists uses snapshot session for read operations.""" driver = MagicMock() @@ -128,7 +106,7 @@ def test_exists_uses_snapshot_session() -> None: config.provide_session.return_value = cm store = SpannerSyncStore(config) - result = store._exists("s1") + result = store._exists("s1") # pyright: ignore config.provide_session.assert_called_once_with() assert result is True diff --git a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py index 1969ee7d8..f459d5848 100644 --- a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py +++ b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py @@ -1,7 +1,6 @@ """Spanner load_from_arrow mutations transport (insert_or_update).""" from typing import Any, cast -from unittest.mock import MagicMock import pyarrow as pa import pytest @@ -23,8 +22,6 @@ def __init__(self) -> None: self.insert_or_update_calls: list[tuple[str, list[str], list[list[Any]]]] = [] self.execute_update_calls: list[str] = [] self.committed = None - self._database = MagicMock() - self._database.execute_partitioned_dml.return_value = 0 def insert_or_update(self, table: str, columns: Any, values: Any) -> None: self.insert_or_update_calls.append((table, list(columns), [list(v) for v in values])) @@ -73,9 +70,8 @@ def test_load_from_arrow_overwrite_deletes_then_mutates(mutations_driver: Spanne mutations_driver.load_from_arrow("users", arrow_table, overwrite=True) - txn._database.execute_partitioned_dml.assert_called_once() - sql = txn._database.execute_partitioned_dml.call_args[0][0] - assert "DELETE FROM users WHERE TRUE" in sql + assert txn.execute_update_calls + assert "DELETE FROM users WHERE TRUE" in txn.execute_update_calls[0] assert len(txn.insert_or_update_calls) == 1 diff --git a/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py b/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py deleted file mode 100644 index 1c45d95d3..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_arrow_overwrite.py +++ /dev/null @@ -1,35 +0,0 @@ -"""Unit tests for load_from_arrow(overwrite=True) using Partitioned DML.""" - -from unittest.mock import MagicMock, patch - -import pyarrow as pa - -from sqlspec.adapters.spanner.driver import SpannerSyncDriver - -CAPABILITIES = { - "arrow_export_enabled": True, - "arrow_import_enabled": True, - "parquet_export_enabled": True, - "parquet_import_enabled": True, - "partition_strategies": ["fixed"], -} - - -def test_load_from_arrow_overwrite_uses_partitioned_dml() -> None: - """Verify load_from_arrow(overwrite=True) calls execute_partitioned_dml for truncation.""" - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 1000 - - mock_connection = MagicMock() - mock_connection._session._database = mock_db - - driver = SpannerSyncDriver(connection=mock_connection, driver_features={"storage_capabilities": CAPABILITIES}) - - arrow_table = pa.table({"id": [1, 2], "name": ["a", "b"]}) - - with patch.object(SpannerSyncDriver, "_arrow_table_to_rows", return_value=(["id", "name"], [])): - driver.load_from_arrow("users", arrow_table, overwrite=True) - - mock_db.execute_partitioned_dml.assert_called_once() - sql = mock_db.execute_partitioned_dml.call_args[0][0] - assert "DELETE FROM users WHERE TRUE" in sql diff --git a/tests/unit/adapters/test_spanner/test_spanner_batch_write.py b/tests/unit/adapters/test_spanner/test_spanner_batch_write.py deleted file mode 100644 index 6e41eaf66..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_batch_write.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Unit tests for load_from_arrow with enable_batch_write_api.""" - -from unittest.mock import MagicMock - -import pyarrow as pa -import pytest - -from sqlspec.adapters.spanner.driver import SpannerSyncDriver -from sqlspec.exceptions import SQLConversionError - -CAPABILITIES = { - "arrow_export_enabled": True, - "arrow_import_enabled": True, - "parquet_export_enabled": True, - "parquet_import_enabled": True, - "partition_strategies": ["fixed"], -} - - -def test_batch_write_succeeds_without_transaction() -> None: - """Verify load_from_arrow with enable_batch_write_api succeeds on non-transaction connection.""" - mock_db = MagicMock() - mock_mg = MagicMock() - mock_db.mutation_groups.return_value.__enter__.return_value = mock_mg - mock_group = MagicMock() - mock_mg.group.return_value = mock_group - mock_response = MagicMock() - mock_response.status = None - mock_mg.batch_write.return_value = [mock_response] - - mock_snapshot = MagicMock() - mock_snapshot._session._database = mock_db - - driver = SpannerSyncDriver( - connection=mock_snapshot, driver_features={"storage_capabilities": CAPABILITIES, "enable_batch_write_api": True} - ) - - arrow_table = pa.table({"id": [1, 2], "name": ["alice", "bob"]}) - - job = driver.load_from_arrow("users", arrow_table) - assert job.telemetry["rows_processed"] == 2 - mock_db.mutation_groups.assert_called_once() - mock_mg.batch_write.assert_called_once() - mock_group.insert_or_update.assert_called_once() - - -def test_standard_insert_requires_transaction() -> None: - """Verify load_from_arrow without enable_batch_write_api still requires a SpannerTransaction.""" - mock_snapshot = MagicMock() - - driver = SpannerSyncDriver( - connection=mock_snapshot, - driver_features={"storage_capabilities": CAPABILITIES, "enable_batch_write_api": False}, - ) - - arrow_table = pa.table({"id": [1, 2], "name": ["alice", "bob"]}) - - with pytest.raises(SQLConversionError, match=r"Arrow import requires a Transaction context\."): - driver.load_from_arrow("users", arrow_table) diff --git a/tests/unit/adapters/test_spanner/test_spanner_last_statement.py b/tests/unit/adapters/test_spanner/test_spanner_last_statement.py deleted file mode 100644 index c0dc62d0a..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_last_statement.py +++ /dev/null @@ -1,58 +0,0 @@ -"""Unit tests for Spanner last_statement execution and commit behavior.""" - -from datetime import datetime, timezone -from unittest.mock import MagicMock - -from sqlspec.adapters.spanner.config import SpannerConnectionContext, SpannerSyncConfig -from sqlspec.adapters.spanner.driver import SpannerSyncDriver - - -def test_execute_passes_last_statement_to_writer() -> None: - """Verify that driver.execute forwards last_statement=True to writer.execute_update.""" - mock_cursor = MagicMock() - mock_cursor.execute_update.return_value = 1 - mock_cursor.committed = None - - driver = SpannerSyncDriver(connection=mock_cursor) - driver.execute("UPDATE t SET x = 1 WHERE id = 'a'", last_statement=True) - - mock_cursor.execute_update.assert_called_once() - _, kwargs = mock_cursor.execute_update.call_args - assert kwargs.get("last_statement") is True - - -def test_driver_commit_noop_when_transaction_already_committed() -> None: - """Verify that driver.commit is a no-op when writer.committed is set.""" - mock_cursor = MagicMock() - mock_cursor.committed = datetime.now(timezone.utc) - mock_cursor.commit = MagicMock() - - driver = SpannerSyncDriver(connection=mock_cursor) - driver.commit() - - mock_cursor.commit.assert_not_called() - - -def test_connection_context_exit_noop_when_already_committed() -> None: - """Verify that SpannerConnectionContext.__exit__ skips commit when txn.committed is set.""" - mock_txn = MagicMock() - mock_txn._transaction_id = b"tx1" - mock_txn.committed = datetime.now(timezone.utc) - mock_txn.commit = MagicMock() - - mock_session = MagicMock() - mock_session.transaction.return_value = mock_txn - - mock_db = MagicMock() - mock_db.sessions_manager.put_session = MagicMock() - - config = MagicMock(spec=SpannerSyncConfig) - config.get_database.return_value = mock_db - - ctx = SpannerConnectionContext(config, transaction=True) - ctx._session = mock_session - ctx._connection = mock_txn - - ctx.__exit__(None, None, None) - - mock_txn.commit.assert_not_called() diff --git a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py b/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py deleted file mode 100644 index 4f905bf6a..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_partitioned_dml.py +++ /dev/null @@ -1,99 +0,0 @@ -"""Unit tests for Spanner Partitioned DML execution on driver.""" - -from unittest.mock import MagicMock - -import pytest -from google.cloud.spanner_v1.types.type import TypeCode - -from sqlspec.adapters.spanner.core import default_statement_config -from sqlspec.adapters.spanner.driver import SpannerSyncDriver -from sqlspec.core import SQL -from sqlspec.exceptions import SQLConversionError - - -def test_driver_execute_partitioned_dml() -> None: - """Verify that driver._execute_partitioned_dml delegates to database.execute_partitioned_dml.""" - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 42 - - mock_connection = MagicMock() - mock_connection._session._database = mock_db - - driver = SpannerSyncDriver( - connection=mock_connection, statement_config=default_statement_config, driver_features={} - ) - - rows = driver._execute_partitioned_dml("DELETE FROM large_table WHERE active = FALSE") - assert rows == 42 - mock_db.execute_partitioned_dml.assert_called_once() - sql = mock_db.execute_partitioned_dml.call_args[0][0] - assert "DELETE FROM large_table WHERE active = FALSE" in sql - - -def test_driver_execute_partitioned_dml_with_parameters() -> None: - """Verify parameters and types are coerced and passed to execute_partitioned_dml.""" - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 10 - - mock_connection = MagicMock() - mock_connection._session._database = mock_db - - driver = SpannerSyncDriver( - connection=mock_connection, statement_config=default_statement_config, driver_features={} - ) - - rows = driver._execute_partitioned_dml( - "UPDATE large_table SET status = :status WHERE threshold > :limit", {"status": "archived", "limit": 100} - ) - assert rows == 10 - mock_db.execute_partitioned_dml.assert_called_once() - _, kwargs = mock_db.execute_partitioned_dml.call_args - assert kwargs["params"] == {"status": "archived", "limit": 100} - assert "status" in kwargs["param_types"] - assert kwargs["param_types"]["status"].code == TypeCode.STRING - assert "limit" in kwargs["param_types"] - assert kwargs["param_types"]["limit"].code == TypeCode.INT64 - - -def test_driver_execute_partitioned_dml_with_sql_object_and_options() -> None: - """Verify executing partitioned DML with SQL object, query options, and request options.""" - mock_db = MagicMock() - mock_db.execute_partitioned_dml.return_value = 50 - - mock_connection = MagicMock() - mock_connection._session._database = mock_db - - driver = SpannerSyncDriver( - connection=mock_connection, statement_config=default_statement_config, driver_features={} - ) - - statement = SQL("DELETE FROM large_table WHERE expired = TRUE", statement_config=default_statement_config) - mock_query_options = MagicMock() - mock_request_options = MagicMock() - - rows = driver._execute_partitioned_dml( - statement, - query_options=mock_query_options, - request_options=mock_request_options, - exclude_txn_from_change_streams=True, - ) - assert rows == 50 - mock_db.execute_partitioned_dml.assert_called_once() - _, kwargs = mock_db.execute_partitioned_dml.call_args - assert kwargs["query_options"] is mock_query_options - assert kwargs["request_options"] is mock_request_options - assert kwargs["exclude_txn_from_change_streams"] is True - - -def test_driver_execute_partitioned_dml_no_database_raises() -> None: - """Verify error raised when database cannot be resolved.""" - mock_connection = MagicMock() - mock_connection._session = None - mock_connection._database = None - - driver = SpannerSyncDriver( - connection=mock_connection, statement_config=default_statement_config, driver_features={} - ) - - with pytest.raises(SQLConversionError, match="Could not resolve Spanner database"): - driver._execute_partitioned_dml("DELETE FROM large_table WHERE TRUE") diff --git a/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py b/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py deleted file mode 100644 index 2df4e3b7c..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_pinging_pool.py +++ /dev/null @@ -1,51 +0,0 @@ -"""Unit tests for Spanner PingingPool fallback and configuration.""" - -from google.cloud.spanner_v1.pool import PingingPool - -from sqlspec.adapters.spanner.config import SpannerSyncConfig - - -def test_disable_multiplexed_sessions_defaults_to_pinging_pool() -> None: - """Verify that disabling multiplexed sessions defaults to PingingPool with 1800s interval.""" - config = SpannerSyncConfig( - connection_config={ - "project": "test-project", - "instance_id": "test-instance", - "database_id": "test-db", - "enable_multiplexed_sessions": False, - } - ) - assert config.connection_config.get("pool_type") is PingingPool - assert config.connection_config.get("ping_interval") == 1800 - - pool = config.provide_pool() - assert isinstance(pool, PingingPool) - assert pool._delta.total_seconds() == 1800 - - -def test_pinging_pool_custom_ping_interval() -> None: - """Verify that custom ping_interval is respected when configuring PingingPool.""" - config = SpannerSyncConfig( - connection_config={ - "project": "test-project", - "instance_id": "test-instance", - "database_id": "test-db", - "enable_multiplexed_sessions": False, - "ping_interval": 900, - } - ) - assert config.connection_config.get("ping_interval") == 900 - - pool = config.provide_pool() - assert isinstance(pool, PingingPool) - assert pool._delta.total_seconds() == 900 - - -def test_provide_pool_fallback_defaults_to_pinging_pool() -> None: - """Verify that calling provide_pool directly falls back to PingingPool with default ping_interval.""" - config = SpannerSyncConfig( - connection_config={"project": "test-project", "instance_id": "test-instance", "database_id": "test-db"} - ) - pool = config.provide_pool() - assert isinstance(pool, PingingPool) - assert pool._delta.total_seconds() == 1800 diff --git a/tests/unit/adapters/test_spanner/test_spanner_pool.py b/tests/unit/adapters/test_spanner/test_spanner_pool.py deleted file mode 100644 index 8dfbb52c3..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_pool.py +++ /dev/null @@ -1,71 +0,0 @@ -"""Unit tests for Spanner session pool configuration and multiplexed pooling.""" - -from unittest.mock import MagicMock, patch - -from google.cloud.spanner_v1.pool import BurstyPool - -from sqlspec.adapters.spanner.config import SpannerSyncConfig - - -def test_multiplexed_session_pool_default() -> None: - """Verify that default SpannerSyncConfig uses multiplexed pooling and omits pool from database args.""" - config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) - assert config.connection_config.get("pool_type") is None - - mock_db = MagicMock() - mock_instance = MagicMock() - mock_instance.database.return_value = mock_db - mock_client = MagicMock() - mock_client.instance.return_value = mock_instance - - with patch.object(config, "_get_client", return_value=mock_client): - db = config.get_database() - - assert db is mock_db - mock_instance.database.assert_called_once() - _, kwargs = mock_instance.database.call_args - assert "pool" not in kwargs - - -def test_explicit_pool_type_preserved() -> None: - """Verify that explicit pool_type in connection_config is instantiated and forwarded.""" - config = SpannerSyncConfig( - connection_config={"project": "p", "instance_id": "i", "database_id": "d", "pool_type": BurstyPool} - ) - assert config.connection_config.get("pool_type") is BurstyPool - - mock_db = MagicMock() - mock_instance = MagicMock() - mock_instance.database.return_value = mock_db - mock_client = MagicMock() - mock_client.instance.return_value = mock_instance - - with patch.object(config, "_get_client", return_value=mock_client): - db = config.get_database() - - assert db is mock_db - mock_instance.database.assert_called_once() - _, kwargs = mock_instance.database.call_args - assert "pool" in kwargs - assert isinstance(kwargs["pool"], BurstyPool) - - -def test_disable_multiplexed_sessions_uses_legacy_pool() -> None: - """Verify that enable_multiplexed_sessions=False constructs an explicit session pool.""" - config = SpannerSyncConfig( - connection_config={"project": "p", "instance_id": "i", "database_id": "d", "enable_multiplexed_sessions": False} - ) - - mock_db = MagicMock() - mock_instance = MagicMock() - mock_instance.database.return_value = mock_db - mock_client = MagicMock() - mock_client.instance.return_value = mock_instance - - with patch.object(config, "_get_client", return_value=mock_client): - db = config.get_database() - - assert db is mock_db - mock_instance.database.assert_called_once() - _, kwargs = mock_instance.database.call_args - assert "pool" in kwargs diff --git a/tests/unit/adapters/test_spanner/test_spanner_query_options.py b/tests/unit/adapters/test_spanner/test_spanner_query_options.py deleted file mode 100644 index 665875388..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_query_options.py +++ /dev/null @@ -1,101 +0,0 @@ -"""Unit tests for Spanner QueryOptions forwarding.""" - -from unittest.mock import MagicMock - -from sqlspec.adapters.spanner.config import SpannerSyncConfig -from sqlspec.adapters.spanner.core import default_statement_config -from sqlspec.adapters.spanner.driver import SpannerSyncDriver - - -def test_driver_execute_with_statement_query_options() -> None: - """Verify driver.execute passes query_options to execute_sql.""" - mock_cursor = MagicMock() - mock_result_set = MagicMock() - mock_field = MagicMock() - mock_field.name = "v" - mock_field.type_.code = 1 - mock_result_set.metadata.row_type.fields = [mock_field] - mock_result_set.__iter__.return_value = [[1]] - mock_cursor.execute_sql.return_value = mock_result_set - - driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) - - query_opts = {"optimizer_version": "6", "optimizer_statistics_package": "latest"} - driver.execute("SELECT 1", query_options=query_opts) - - mock_cursor.execute_sql.assert_called_once() - _, kwargs = mock_cursor.execute_sql.call_args - assert kwargs.get("query_options") == query_opts - - -def test_driver_feature_query_options() -> None: - """Verify driver-level query_options are passed to execute_sql.""" - mock_cursor = MagicMock() - mock_result_set = MagicMock() - mock_field = MagicMock() - mock_field.name = "v" - mock_field.type_.code = 1 - mock_result_set.metadata.row_type.fields = [mock_field] - mock_result_set.__iter__.return_value = [[1]] - mock_cursor.execute_sql.return_value = mock_result_set - - query_opts = {"optimizer_version": "latest"} - driver = SpannerSyncDriver( - connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} - ) - - driver.execute("SELECT 1") - - mock_cursor.execute_sql.assert_called_once() - _, kwargs = mock_cursor.execute_sql.call_args - assert kwargs.get("query_options") == query_opts - - -def test_config_provide_session_query_options() -> None: - """Verify config.provide_session forwards query_options to driver_features.""" - config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) - query_opts = {"optimizer_version": "5"} - features = config._session_driver_features( - request_options=None, directed_read_options=None, query_options=query_opts, retry=None, timeout=None - ) - assert features["query_options"] == query_opts - - -def test_driver_execute_many_omits_query_options() -> None: - """Verify execute_many does not forward query_options to batch_update.""" - mock_cursor = MagicMock() - mock_cursor.batch_update.return_value = (None, [1, 1]) - - query_opts = {"optimizer_version": "latest"} - driver = SpannerSyncDriver( - connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} - ) - - driver.execute_many("INSERT INTO t (id) VALUES (:id)", [{"id": 1}, {"id": 2}], query_options=query_opts) - - mock_cursor.batch_update.assert_called_once() - _, kwargs = mock_cursor.batch_update.call_args - assert "query_options" not in kwargs - - -def test_driver_select_stream_query_options() -> None: - """Verify select_stream passes query_options to execute_sql.""" - mock_cursor = MagicMock() - mock_result_set = MagicMock() - mock_field = MagicMock() - mock_field.name = "v" - mock_field.type_.code = 1 - mock_result_set.metadata.row_type.fields = [mock_field] - mock_result_set.__iter__.return_value = [[1], [2]] - mock_cursor.execute_sql.return_value = mock_result_set - - query_opts = {"optimizer_version": "6"} - driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) - - stream = driver.select_stream("SELECT 1", query_options=query_opts) - rows = list(stream) if stream else [] - assert len(rows) == 2 - - mock_cursor.execute_sql.assert_called_once() - _, kwargs = mock_cursor.execute_sql.call_args - assert kwargs.get("query_options") == query_opts diff --git a/tests/unit/adapters/test_spanner/test_spanner_request_options.py b/tests/unit/adapters/test_spanner/test_spanner_request_options.py deleted file mode 100644 index 1aff60507..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_request_options.py +++ /dev/null @@ -1,105 +0,0 @@ -"""Unit tests for Spanner RequestOptions and DirectedReadOptions propagation.""" - -from unittest.mock import MagicMock - -from sqlspec.adapters.spanner.core import default_statement_config -from sqlspec.adapters.spanner.driver import SpannerSyncDriver - - -def test_execute_select_forwards_request_and_directed_read_options() -> None: - """Verify execute for SELECT forwards request_options and directed_read_options.""" - mock_cursor = MagicMock() - mock_result_set = MagicMock() - mock_field = MagicMock() - mock_field.name = "v" - mock_field.type_.code = 1 - mock_result_set.metadata.row_type.fields = [mock_field] - mock_result_set.__iter__.return_value = [[1]] - mock_cursor.execute_sql.return_value = mock_result_set - - driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) - - req_opts = {"request_tag": "select-tag", "priority": 1} - directed_read = MagicMock() - driver.execute("SELECT 1", request_options=req_opts, directed_read_options=directed_read) - - mock_cursor.execute_sql.assert_called_once() - _, kwargs = mock_cursor.execute_sql.call_args - assert kwargs.get("request_options") == req_opts - assert kwargs.get("directed_read_options") is directed_read - - -def test_execute_update_omits_directed_read_options() -> None: - """Verify execute for UPDATE forwards request_options but strips directed_read_options.""" - mock_cursor = MagicMock() - mock_cursor.execute_update.return_value = 1 - - driver = SpannerSyncDriver( - connection=mock_cursor, - statement_config=default_statement_config, - driver_features={"directed_read_options": MagicMock()}, - ) - - req_opts = {"request_tag": "update-tag", "transaction_tag": "tx-tag"} - per_call_directed = MagicMock() - driver.execute("UPDATE t SET x = 1", request_options=req_opts, directed_read_options=per_call_directed) - - mock_cursor.execute_update.assert_called_once() - _, kwargs = mock_cursor.execute_update.call_args - assert kwargs.get("request_options") == req_opts - assert "directed_read_options" not in kwargs - - -def test_execute_many_forwards_request_options_and_omits_read_options() -> None: - """Verify execute_many forwards request_options and omits directed_read_options and query_options.""" - mock_cursor = MagicMock() - mock_cursor.batch_update.return_value = (None, [1, 1]) - - driver = SpannerSyncDriver( - connection=mock_cursor, - statement_config=default_statement_config, - driver_features={"query_options": {"optimizer_version": "6"}, "directed_read_options": MagicMock()}, - ) - - req_opts = {"request_tag": "batch-tag", "priority": 2} - driver.execute_many( - "INSERT INTO t (id) VALUES (:id)", - [{"id": 1}, {"id": 2}], - request_options=req_opts, - directed_read_options=MagicMock(), - query_options={"optimizer_version": "5"}, - ) - - mock_cursor.batch_update.assert_called_once() - _, kwargs = mock_cursor.batch_update.call_args - assert kwargs.get("request_options") == req_opts - assert "directed_read_options" not in kwargs - assert "query_options" not in kwargs - - -def test_execute_script_forwards_options_appropriately() -> None: - """Verify execute_script separates read and write options across script statements.""" - mock_cursor = MagicMock() - mock_result_set = MagicMock() - mock_cursor.execute_sql.return_value = mock_result_set - mock_cursor.execute_update.return_value = 1 - - req_opts = {"request_tag": "script-tag"} - directed_read = MagicMock() - driver = SpannerSyncDriver( - connection=mock_cursor, - statement_config=default_statement_config, - driver_features={"request_options": req_opts, "directed_read_options": directed_read}, - ) - - driver.execute_script("SELECT 1; UPDATE t SET x = 1;") - - mock_cursor.execute_sql.assert_called_once() - _, read_kwargs = mock_cursor.execute_sql.call_args - assert read_kwargs.get("request_options") == req_opts - assert read_kwargs.get("directed_read_options") is directed_read - - mock_cursor.execute_update.assert_called_once() - _, write_kwargs = mock_cursor.execute_update.call_args - assert write_kwargs.get("request_options") == req_opts - assert "directed_read_options" not in write_kwargs diff --git a/tests/unit/adapters/test_spanner/test_spanner_stores.py b/tests/unit/adapters/test_spanner/test_spanner_stores.py deleted file mode 100644 index 63d960441..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_stores.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Unit tests for ADK and Litestar store session routing.""" - -from typing import Any -from unittest.mock import MagicMock - -from sqlspec.adapters.spanner._typing import spanner_param_types as param_types -from sqlspec.adapters.spanner.adk import SpannerSyncADKMemoryStore, SpannerSyncADKStore -from sqlspec.adapters.spanner.config import SpannerSyncConfig -from sqlspec.adapters.spanner.driver import SpannerSyncDriver -from sqlspec.adapters.spanner.litestar import SpannerSyncStore -from sqlspec.adapters.spanner.type_converter import bytes_to_spanner, spanner_json -from sqlspec.core import TypedParameter - - -def _context_manager_yielding(value: Any) -> Any: - class _Ctx: - def __enter__(self) -> Any: - return value - - def __exit__(self, *_: Any) -> None: - pass - - return _Ctx() - - -def test_adk_store_run_write_routes_through_provide_session() -> None: - """Verify that SpannerSyncADKStore._run_write executes via config.provide_session(transaction=True).""" - config = MagicMock(spec=SpannerSyncConfig) - executed_statements: list[tuple[str, Any]] = [] - - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_driver.execute.side_effect = lambda sql, params=None, param_types=None, **kw: executed_statements.append(( - sql, - params, - )) - config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) - - store = SpannerSyncADKStore(config=config) - statements = [ - ("INSERT INTO t (id) VALUES (@id)", {"id": "1"}, {"id": MagicMock()}), - ("INSERT INTO t (id) VALUES (@id)", {"id": "2"}, {"id": MagicMock()}), - ] - store._run_write(statements) - - config.provide_session.assert_called_once_with(transaction=True) - assert len(executed_statements) == 2 - - -def test_litestar_store_writes_route_through_provide_session() -> None: - """Verify that SpannerSyncStore write operations execute via config.provide_session(transaction=True).""" - config = MagicMock(spec=SpannerSyncConfig) - config.extension_config = {"litestar": {"session_table": "sessions"}} - executed_sqls: list[str] = [] - - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_result = MagicMock() - mock_result.rowcount = 1 - mock_result.rows_affected = 1 - - def mock_execute(sql: Any, *a: Any, **kw: Any) -> Any: - executed_sqls.append(str(sql)) - return mock_result - - mock_driver.execute.side_effect = mock_execute - config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) - - store = SpannerSyncStore(config=config) - - store._set("session_1", b"payload", expires_in=3600) - assert config.provide_session.call_count == 1 - config.provide_session.assert_called_with(transaction=True) - - store._delete("session_1") - assert config.provide_session.call_count == 2 - config.provide_session.assert_called_with(transaction=True) - - store._delete_all() - assert config.provide_session.call_count == 3 - config.provide_session.assert_called_with(transaction=True) - - expired_count = store._delete_expired() - assert config.provide_session.call_count == 4 - config.provide_session.assert_called_with(transaction=True) - assert expired_count == 1 - - -def test_litestar_store_single_base64_roundtrip() -> None: - """Verify SpannerSyncStore passes raw bytes to driver.execute and decodes wire bytes once on _get.""" - config = MagicMock(spec=SpannerSyncConfig) - config.extension_config = {"litestar": {"session_table": "sessions"}} - captured_params: list[dict[str, Any]] = [] - - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_result = MagicMock() - mock_result.rows_affected = 1 - - def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: - if isinstance(params, dict): - captured_params.append(params) - return mock_result - - mock_driver.execute.side_effect = mock_execute - mock_driver.select_one_or_none.return_value = {"data": bytes_to_spanner(b"raw-payload"), "expires_at": None} - config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) - - store = SpannerSyncStore(config=config) - store._set("session_1", b"raw-payload", expires_in=None) - config.provide_session.assert_called_once_with(transaction=True) - - assert len(captured_params) == 1 - assert captured_params[0]["data"] == b"raw-payload" - assert isinstance(captured_params[0]["expires_at"], TypedParameter) - assert captured_params[0]["expires_at"].value is None - - fetched = store._get("session_1") - assert fetched == b"raw-payload" - - -def test_adk_memory_store_write_and_decode_json() -> None: - """Verify SpannerSyncADKMemoryStore prepares JSON/null write params and unwraps JsonObject.""" - config = MagicMock(spec=SpannerSyncConfig) - config.extension_config = {"adk": {"enable_memory": True}} - captured_params: list[dict[str, Any]] = [] - - mock_driver = MagicMock(spec=SpannerSyncDriver) - mock_result = MagicMock() - mock_result.rows_affected = 5 - - def mock_execute(_sql: Any, params: Any = None, *_a: Any, **_kw: Any) -> Any: - if isinstance(params, dict): - captured_params.append(params) - return mock_result - - mock_driver.execute.side_effect = mock_execute - config.provide_session.side_effect = lambda *a, **kw: _context_manager_yielding(mock_driver) - - store = SpannerSyncADKMemoryStore(config=config) - store._run_write([ - ( - "INSERT INTO adk_memory VALUES (@content_json, @metadata_json, @owner_id)", - {"content_json": '{"text":"hi"}', "metadata_json": None, "owner_id": None}, - {"content_json": param_types.JSON, "metadata_json": param_types.JSON, "owner_id": param_types.STRING}, - ) - ]) - config.provide_session.assert_called_once_with(transaction=True) - - assert len(captured_params) == 1 - assert captured_params[0]["content_json"] == {"text": "hi"} - assert isinstance(captured_params[0]["metadata_json"], TypedParameter) - assert captured_params[0]["metadata_json"].original_type is dict - assert isinstance(captured_params[0]["owner_id"], TypedParameter) - assert captured_params[0]["owner_id"].original_type is str - - deleted = store._execute_update( - "DELETE FROM adk_memory WHERE session_id = @session_id", - {"session_id": "s1"}, - {"session_id": param_types.STRING}, - ) - assert config.provide_session.call_count == 2 - config.provide_session.assert_called_with(transaction=True) - assert deleted == 5 - - decoded = store._decode_json(spanner_json({"k": "v"})) - assert decoded == {"k": "v"} diff --git a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py index a68be300e..673bcc792 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py +++ b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py @@ -1,6 +1,5 @@ -"""Unit tests for Decimal and INTERVAL parameter type inference.""" +"""Unit tests for Decimal and JSON parameter type inference.""" -from datetime import timedelta from decimal import Decimal from google.cloud.spanner_v1.types.type import TypeCode @@ -26,25 +25,8 @@ def test_infer_decimal_array_param_types() -> None: assert types["prices"].array_element_type.code == TypeCode.NUMERIC -def test_infer_timedelta_param_types() -> None: - """Verify that timedelta parameters infer as INTERVAL.""" - params = {"duration": timedelta(days=1, hours=2)} - types = infer_spanner_param_types(params) - assert "duration" in types - assert types["duration"].code == TypeCode.INTERVAL - - -def test_null_timedelta_param_types() -> None: - """Verify that null TypedParameter with timedelta resolves to INTERVAL.""" - params = {"duration": TypedParameter(None, timedelta)} - types = infer_spanner_param_types(params) - assert "duration" in types - assert types["duration"].code == TypeCode.INTERVAL - - def test_null_json_param_types() -> None: - """Verify that null TypedParameter with dict or JSON resolves to JSON.""" - params = {"meta": TypedParameter(None, dict), "payload": TypedParameter(None, "JSON")} + """Verify that null TypedParameter with dict resolves to JSON.""" + params = {"meta": TypedParameter(None, dict)} types = infer_spanner_param_types(params) assert types["meta"].code == TypeCode.JSON - assert types["payload"].code == TypeCode.JSON diff --git a/tests/unit/adapters/test_spanner/test_spanner_vector.py b/tests/unit/adapters/test_spanner/test_spanner_vector.py deleted file mode 100644 index ab0390adb..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_vector.py +++ /dev/null @@ -1,48 +0,0 @@ -"""Unit tests for Spanner FLOAT32 and ARRAY vector parameter typing.""" - -from google.cloud.spanner_v1.types.type import TypeCode - -from sqlspec.adapters.spanner.type_converter import coerce_params_for_spanner, infer_spanner_param_types -from sqlspec.core import TypedParameter - - -def test_infer_vector_param_types_array_float32() -> None: - """Verify that TypedParameter with ARRAY infers as Array(FLOAT32).""" - params = {"embedding": TypedParameter([0.1, 0.2, 0.3], "ARRAY")} - types = infer_spanner_param_types(params) - assert "embedding" in types - assert types["embedding"].code == TypeCode.ARRAY - assert types["embedding"].array_element_type.code == TypeCode.FLOAT32 - - -def test_infer_vector_param_types_vector_alias() -> None: - """Verify that TypedParameter with VECTOR infers as Array(FLOAT32).""" - params = {"embedding": TypedParameter([0.1, 0.2, 0.3], "VECTOR")} - types = infer_spanner_param_types(params) - assert "embedding" in types - assert types["embedding"].code == TypeCode.ARRAY - assert types["embedding"].array_element_type.code == TypeCode.FLOAT32 - - -def test_infer_scalar_float32() -> None: - """Verify that TypedParameter with FLOAT32 infers as FLOAT32.""" - params = {"score": TypedParameter(1.25, "FLOAT32")} - types = infer_spanner_param_types(params) - assert "score" in types - assert types["score"].code == TypeCode.FLOAT32 - - -def test_coerce_vector_params() -> None: - """Verify that TypedParameter vector is unwrapped into a list of floats.""" - params = {"embedding": TypedParameter((0.1, 0.2, 0.3), "ARRAY")} - coerced = coerce_params_for_spanner(params) - assert coerced is not None - assert coerced["embedding"] == [0.1, 0.2, 0.3] - - -def test_null_float32_param_type() -> None: - """Verify that NULL TypedParameter with FLOAT32 resolves to param_types.FLOAT32.""" - params = {"score": TypedParameter(None, "FLOAT32")} - types = infer_spanner_param_types(params) - assert "score" in types - assert types["score"].code == TypeCode.FLOAT32 From 68dee6a50741f01d30f42ff6b913f1abbf367506 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 20:55:16 +0000 Subject: [PATCH 10/15] fix(spanner): retain JSON callback contract and cover boolean arrays --- sqlspec/adapters/spanner/core.py | 38 +++---------- .../test_spanner/test_spanner_json.py | 55 ------------------- .../test_spanner_type_inference.py | 6 ++ 3 files changed, 15 insertions(+), 84 deletions(-) delete mode 100644 tests/unit/adapters/test_spanner/test_spanner_json.py diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index 040b3ae3e..d56dbf572 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -310,39 +310,19 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec return _create_spanner_error(error, SQLSpecError, "error") -def _unwrap_spanner_json_object(val: Any) -> Any: - """Recursively unwrap Spanner JsonObject instances into native Python primitives.""" - if isinstance(val, JsonObject): - if getattr(val, "_is_null", False): - return None - if getattr(val, "_is_array", False): - array_val = getattr(val, "_array_value", None) - return [_unwrap_spanner_json_object(item) for item in array_val] if array_val is not None else [] - if getattr(val, "_is_scalar_value", False): - return getattr(val, "_simple_value", None) - return {k: _unwrap_spanner_json_object(v) for k, v in cast("dict[str, Any]", val).items()} - if isinstance(val, dict): - return {k: _unwrap_spanner_json_object(v) for k, v in val.items()} - if isinstance(val, (list, tuple)): - return [_unwrap_spanner_json_object(item) for item in val] - return val - - def _convert_json_row_value(value: Any, *, json_deserializer: "Callable[[str], Any]") -> Any: """Convert a native Spanner JSON cell using the configured deserializer.""" if isinstance(value, JsonObject): - if json_deserializer is from_json: - return _unwrap_spanner_json_object(value) - try: - return json_deserializer(cast("Any", value).serialize()) - except (TypeError, ValueError): - return value + json_value = cast("Any", value).serialize() elif isinstance(value, str): - try: - return json_deserializer(value) - except (TypeError, ValueError): - return value - return value + json_value = value + else: + return value + + try: + return json_deserializer(json_value) + except (TypeError, ValueError): + return value def _create_spanner_error(error: Any, error_class: type[SQLSpecError], description: str) -> SQLSpecError: diff --git a/tests/unit/adapters/test_spanner/test_spanner_json.py b/tests/unit/adapters/test_spanner/test_spanner_json.py deleted file mode 100644 index 8f2dd71e1..000000000 --- a/tests/unit/adapters/test_spanner/test_spanner_json.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Unit tests for Spanner JsonObject direct unwrapping optimization.""" - -from typing import Any, cast -from unittest.mock import MagicMock - -from google.cloud.spanner_v1.data_types import JsonObject - -from sqlspec.adapters.spanner.core import _convert_json_row_value -from sqlspec.utils.serializers import from_json - - -class MonitoredJsonObject(JsonObject): - """JsonObject subclass that tracks calls to serialize().""" - - def __init__(self, *args: Any, **kwargs: Any) -> None: - cast("Any", super()).__init__(*args, **kwargs) - self.serialize_called = False - - def serialize(self) -> str | None: - self.serialize_called = True - return cast("str | None", cast("Any", super()).serialize()) - - -def test_convert_json_row_value_unwraps_without_serialize() -> None: - """Verify that default deserializer unwraps JsonObject directly without calling serialize().""" - obj = MonitoredJsonObject({"key": "val", "nested": [1, 2]}) - res = _convert_json_row_value(obj, json_deserializer=from_json) - assert res == {"key": "val", "nested": [1, 2]} - assert not obj.serialize_called - - -def test_convert_json_row_value_calls_serialize_for_custom_deserializer() -> None: - """Verify that a custom string deserializer invokes serialize().""" - obj = MonitoredJsonObject({"key": "val"}) - custom_deserializer = MagicMock(return_value={"custom": True}) - res = _convert_json_row_value(obj, json_deserializer=custom_deserializer) - assert res == {"custom": True} - assert obj.serialize_called - custom_deserializer.assert_called_once_with('{"key":"val"}') - - -def test_convert_json_row_value_null_json() -> None: - """Verify that null JsonObject unwraps directly to None.""" - obj = MonitoredJsonObject(None) - res = _convert_json_row_value(obj, json_deserializer=from_json) - assert res is None - assert not obj.serialize_called - - -def test_convert_json_row_value_array_json() -> None: - """Verify that array JsonObject unwraps directly to a list.""" - obj = MonitoredJsonObject([1, 2, 3]) - res = _convert_json_row_value(obj, json_deserializer=from_json) - assert res == [1, 2, 3] - assert not obj.serialize_called diff --git a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py index 673bcc792..3ff2dd354 100644 --- a/tests/unit/adapters/test_spanner/test_spanner_type_inference.py +++ b/tests/unit/adapters/test_spanner/test_spanner_type_inference.py @@ -30,3 +30,9 @@ def test_null_json_param_types() -> None: params = {"meta": TypedParameter(None, dict)} types = infer_spanner_param_types(params) assert types["meta"].code == TypeCode.JSON + + +def test_infer_boolean_array_param_types() -> None: + types = infer_spanner_param_types({"flags": [True, False]}) + assert types["flags"].code == TypeCode.ARRAY + assert types["flags"].array_element_type.code == TypeCode.BOOL From 08466f81bdb16db685d4c6942be856585c371b76 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:36:16 +0000 Subject: [PATCH 11/15] fix: preserve native Spanner execution controls --- docs/changelog.rst | 3 + docs/reference/adapters/spanner.rst | 12 ++ sqlspec/adapters/spanner/config.py | 19 ++- sqlspec/adapters/spanner/core.py | 2 + sqlspec/adapters/spanner/driver.py | 59 +++++++--- .../test_spanner/test_batch_write_api.py | 13 +++ .../test_spanner_last_statement.py | 30 +++++ .../test_spanner_query_options.py | 109 ++++++++++++++++++ 8 files changed, 232 insertions(+), 15 deletions(-) create mode 100644 tests/unit/adapters/test_spanner/test_spanner_last_statement.py create mode 100644 tests/unit/adapters/test_spanner/test_spanner_query_options.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 4aa6a5cb8..43137c029 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -25,6 +25,9 @@ Unreleased and later when the SQLite runtime supports them. Choose a transaction lock mode or set the batch size for Arrow imports. Defaults stay the same. +* Spanner forwards native query options and final-statement hints and + supports opt-in Batch Write from read sessions. + * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. diff --git a/docs/reference/adapters/spanner.rst b/docs/reference/adapters/spanner.rst index 092760e47..b95400ba1 100644 --- a/docs/reference/adapters/spanner.rst +++ b/docs/reference/adapters/spanner.rst @@ -139,3 +139,15 @@ namespace: ``"litestar"``, ``"events"``, or ``"adk"`` as supported by this adapt .. autoclass:: sqlspec.adapters.spanner.adk.SpannerADKRetentionConfig :members: :show-inheritance: + +Native execution controls +------------------------- + +``query_options`` can be configured on the driver, supplied when opening a +session, or overridden per call. They apply to queries and single DML operations; +native batch DML does not accept them. ``last_statement=True`` marks final +transaction DML, including only the final statement of a script. + +Opt-in Arrow Batch Write ingestion works from database-backed read sessions. +Mutation groups commit independently. Arrow overwrite retains transactional +delete-and-insert behavior without partitioned DML or Batch Write. diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index aa475eafd..c4b5f9328 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -173,6 +173,7 @@ class SpannerDriverFeatures(TypedDict): json_deserializer: "NotRequired[Callable[[str], Any]]" retry: "NotRequired[Retry | None]" timeout: "NotRequired[float | None]" + query_options: "NotRequired[ExecuteSqlRequest.QueryOptions | dict[str, Any] | None]" request_options: "NotRequired[RequestOptions | dict[str, Any] | None]" directed_read_options: "NotRequired[DirectedReadOptions | None]" session_labels: "NotRequired[dict[str, str]]" @@ -480,6 +481,7 @@ def provide_session( transaction: "bool" = _DEFAULT_SESSION_TRANSACTION, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -497,6 +499,7 @@ def provide_session( Snapshot (False). request_options: Session-scoped RequestOptions for Spanner statements. directed_read_options: Session-scoped DirectedReadOptions for reads. + query_options: Session-scoped QueryOptions for Spanner statements. retry: Session-scoped retry policy for Spanner statement calls. timeout: Session-scoped timeout for Spanner statement calls. **kwargs: Additional keyword arguments. @@ -514,6 +517,7 @@ def provide_session( driver_features=self._session_driver_features( request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, ), @@ -526,6 +530,7 @@ def provide_write_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -537,6 +542,7 @@ def provide_write_session( transaction=True, request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, **kwargs, @@ -548,6 +554,7 @@ def provide_read_session( statement_config: "StatementConfig | None" = None, request_options: "RequestOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, **kwargs: Any, @@ -563,6 +570,7 @@ def provide_read_session( transaction=False, request_options=request_options, directed_read_options=directed_read_options, + query_options=query_options, retry=retry, timeout=timeout, **kwargs, @@ -573,16 +581,25 @@ def _session_driver_features( *, request_options: "RequestOptions | dict[str, Any] | None", directed_read_options: "DirectedReadOptions | None", + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, retry: "Retry | None", timeout: "float | None", ) -> "dict[str, Any]": - if request_options is None and directed_read_options is None and retry is None and timeout is None: + if ( + request_options is None + and directed_read_options is None + and query_options is None + and retry is None + and timeout is None + ): return self.driver_features driver_features = dict(self.driver_features) if request_options is not None: driver_features["request_options"] = request_options if directed_read_options is not None: driver_features["directed_read_options"] = directed_read_options + if query_options is not None: + driver_features["query_options"] = query_options if retry is not None: driver_features["retry"] = retry if timeout is not None: diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index d56dbf572..b2c863cb0 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -314,6 +314,8 @@ def _convert_json_row_value(value: Any, *, json_deserializer: "Callable[[str], A """Convert a native Spanner JSON cell using the configured deserializer.""" if isinstance(value, JsonObject): json_value = cast("Any", value).serialize() + if json_value is None: + return None elif isinstance(value, str): json_value = value else: diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index 322426006..bb195c112 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -45,6 +45,7 @@ from sqlspec.adapters.spanner._typing import SpannerConnection from sqlspec.adapters.spanner._typing import SpannerDirectedReadOptions as DirectedReadOptions + from sqlspec.adapters.spanner._typing import SpannerExecuteSqlRequest as ExecuteSqlRequest from sqlspec.adapters.spanner._typing import SpannerRequestOptions as RequestOptions from sqlspec.adapters.spanner._typing import SpannerRetry as Retry from sqlspec.builder import QueryBuilder @@ -176,7 +177,7 @@ def dispatch_execute_many(self, cursor: "SpannerConnection", statement: "SQL") - _coerce = self._coerce_params _infer = self._infer_param_types - execute_kwargs = self._execute_kwargs() + execute_kwargs = self._execute_kwargs(for_batch=True) param_types_cache: dict[tuple[tuple[str, type[Any], Any], ...], dict[str, Any]] = {} empty_param_types: dict[str, Any] = {} batch_args: list[tuple[str, dict[str, Any] | None, dict[str, Any]]] = [] @@ -212,7 +213,7 @@ def dispatch_execute_script(self, cursor: "SpannerConnection", statement: "SQL") coerced_params = self._coerce_params(script_params) read_execute_kwargs = self._execute_kwargs(for_read=True) write_execute_kwargs = self._execute_kwargs() - for stmt in statements: + for index, stmt in enumerate(statements): try: parsed = _sqlglot.parse_one(stmt) is_select = isinstance(parsed, _sqlglot_exp.Select) @@ -222,7 +223,12 @@ def dispatch_execute_script(self, cursor: "SpannerConnection", statement: "SQL") raise SQLConversionError(_READ_ONLY_SNAPSHOT_ERROR_MESSAGE) if not is_select and is_transaction: writer = cast("_SpannerWriteProtocol", cursor) - writer.execute_update(stmt, params=coerced_params, param_types=param_types_map, **write_execute_kwargs) + statement_kwargs = write_execute_kwargs + if "last_statement" in write_execute_kwargs and index != len(statements) - 1: + statement_kwargs = { + key: value for key, value in write_execute_kwargs.items() if key != "last_statement" + } + writer.execute_update(stmt, params=coerced_params, param_types=param_types_map, **statement_kwargs) else: _ = list( reader.execute_sql(stmt, params=coerced_params, param_types=param_types_map, **read_execute_kwargs) @@ -449,14 +455,14 @@ def load_from_arrow( columns, records = self._arrow_table_to_rows(arrow_table) if records: - conn = self.connection - if not isinstance(conn, SpannerTransaction): - msg = "Arrow import requires a Transaction context." - raise SQLConversionError(msg) chunks = self._chunk_mutation_rows(columns, records) if self.driver_features.get("enable_batch_write_api") and not overwrite: self._batch_write_mutations(table, columns, chunks) else: + conn = self.connection + if not isinstance(conn, SpannerTransaction): + msg = "Arrow import requires a Transaction context." + raise SQLConversionError(msg) writer = cast("_SpannerWriteProtocol", conn) for chunk in chunks: writer.insert_or_update(table, columns, chunk) @@ -512,36 +518,57 @@ def resolve_rowcount(self, cursor: "SpannerConnection") -> int: """ return 0 - def _execute_kwargs(self, *, for_read: bool = False) -> dict[str, Any]: + def _execute_kwargs(self, *, for_read: bool = False, for_batch: bool = False) -> dict[str, Any]: kwargs: dict[str, Any] = { key: self.driver_features[key] for key in ("retry", "timeout") if key in self.driver_features } request_options = self.driver_features.get("request_options") if request_options is not None: kwargs["request_options"] = request_options - directed_read_options = self.driver_features.get("directed_read_options") - if for_read and directed_read_options is not None: - kwargs["directed_read_options"] = directed_read_options + if not for_batch: + query_options = self.driver_features.get("query_options") + if query_options is not None: + kwargs["query_options"] = query_options + if for_read and not for_batch: + directed_read_options = self.driver_features.get("directed_read_options") + if directed_read_options is not None: + kwargs["directed_read_options"] = directed_read_options pending = self._pending_execute_options if pending is not None: if pending.request_options is not None: kwargs["request_options"] = pending.request_options + if not for_batch and pending.query_options is not None: + kwargs["query_options"] = pending.query_options if pending.retry is not None: kwargs["retry"] = pending.retry if pending.timeout is not None: kwargs["timeout"] = pending.timeout - if for_read and pending.directed_read_options is not None: + if for_read and not for_batch and pending.directed_read_options is not None: kwargs["directed_read_options"] = pending.directed_read_options + if not for_read and pending.last_statement: + kwargs["last_statement"] = True return kwargs def _pop_execute_options(self, kwargs: dict[str, Any]) -> "_PerCallExecuteOptions | None": - if not any(key in kwargs for key in ("request_options", "directed_read_options", "retry", "timeout")): + if not any( + key in kwargs + for key in ( + "request_options", + "query_options", + "directed_read_options", + "retry", + "timeout", + "last_statement", + ) + ): return None return _PerCallExecuteOptions( request_options=kwargs.pop("request_options", None), + query_options=kwargs.pop("query_options", None), directed_read_options=kwargs.pop("directed_read_options", None), retry=kwargs.pop("retry", None), timeout=kwargs.pop("timeout", None), + last_statement=bool(kwargs.pop("last_statement", False)), ) def _chunk_mutation_rows(self, columns: "list[str]", records: "list[tuple[Any, ...]]") -> "list[list[list[Any]]]": @@ -647,20 +674,24 @@ def rollback(self) -> None: ... class _PerCallExecuteOptions: """Per-call Spanner execution options captured for a single dispatch.""" - __slots__ = ("directed_read_options", "request_options", "retry", "timeout") + __slots__ = ("directed_read_options", "last_statement", "query_options", "request_options", "retry", "timeout") def __init__( self, *, request_options: "RequestOptions | dict[str, Any] | None" = None, + query_options: "ExecuteSqlRequest.QueryOptions | dict[str, Any] | None" = None, directed_read_options: "DirectedReadOptions | None" = None, retry: "Retry | None" = None, timeout: "float | None" = None, + last_statement: bool = False, ) -> None: self.request_options = request_options + self.query_options = query_options self.directed_read_options = directed_read_options self.retry = retry self.timeout = timeout + self.last_statement = last_statement class _SpannerSelectStreamSource: diff --git a/tests/unit/adapters/test_spanner/test_batch_write_api.py b/tests/unit/adapters/test_spanner/test_batch_write_api.py index 8a002e02a..cc4c71063 100644 --- a/tests/unit/adapters/test_spanner/test_batch_write_api.py +++ b/tests/unit/adapters/test_spanner/test_batch_write_api.py @@ -117,3 +117,16 @@ def test_batch_write_overwrite_uses_transactional_mutations(batch_write_driver: assert conn.execute_update_calls and "DELETE FROM users WHERE TRUE" in conn.execute_update_calls[0] assert conn.insert_or_update_calls == [("users", ["id"], [[1]])] assert conn.database.mutation_groups_obj.batch_write_calls == 0 + + +def test_batch_write_accepts_database_backed_snapshot() -> None: + from types import SimpleNamespace + + database = _FakeDatabase() + snapshot = SimpleNamespace(_session=_FakeSession(database)) + driver = SpannerSyncDriver( + cast("Any", snapshot), driver_features={"storage_capabilities": CAPABILITIES, "enable_batch_write_api": True} + ) + job = driver.load_from_arrow("users", pa.table({"id": [1]})) + assert job.telemetry["rows_processed"] == 1 + assert database.mutation_groups_obj.batch_write_calls == 1 diff --git a/tests/unit/adapters/test_spanner/test_spanner_last_statement.py b/tests/unit/adapters/test_spanner/test_spanner_last_statement.py new file mode 100644 index 000000000..5caa35627 --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_last_statement.py @@ -0,0 +1,30 @@ +"""Unit tests for Spanner last_statement execution and commit behavior.""" + +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_execute_passes_last_statement_to_writer() -> None: + """Verify that driver.execute forwards last_statement=True to writer.execute_update.""" + mock_cursor = MagicMock() + mock_cursor.execute_update.return_value = 1 + mock_cursor.committed = None + + driver = SpannerSyncDriver(connection=mock_cursor) + driver.execute("UPDATE t SET x = 1 WHERE id = 'a'", last_statement=True) + + mock_cursor.execute_update.assert_called_once() + _, kwargs = mock_cursor.execute_update.call_args + assert kwargs.get("last_statement") is True + + +def test_script_only_marks_final_dml_as_last_statement() -> None: + cursor = MagicMock() + cursor.execute_update.return_value = 1 + driver = SpannerSyncDriver(connection=cursor) + driver.execute_script("UPDATE t SET x = 1; UPDATE t SET x = 2", last_statement=True) + calls = cursor.execute_update.call_args_list + assert len(calls) == 2 + assert "last_statement" not in calls[0].kwargs + assert calls[1].kwargs["last_statement"] is True diff --git a/tests/unit/adapters/test_spanner/test_spanner_query_options.py b/tests/unit/adapters/test_spanner/test_spanner_query_options.py new file mode 100644 index 000000000..e4e973e1a --- /dev/null +++ b/tests/unit/adapters/test_spanner/test_spanner_query_options.py @@ -0,0 +1,109 @@ +"""Unit tests for Spanner QueryOptions forwarding.""" + +from unittest.mock import MagicMock + +from sqlspec.adapters.spanner.config import SpannerSyncConfig +from sqlspec.adapters.spanner.core import default_statement_config +from sqlspec.adapters.spanner.driver import SpannerSyncDriver + + +def test_driver_execute_with_statement_query_options() -> None: + """Verify driver.execute passes query_options to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1]] + mock_cursor.execute_sql.return_value = mock_result_set + + driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) + + query_opts = {"optimizer_version": "6", "optimizer_statistics_package": "latest"} + driver.execute("SELECT 1", query_options=query_opts) + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts + + +def test_driver_feature_query_options() -> None: + """Verify driver-level query_options are passed to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1]] + mock_cursor.execute_sql.return_value = mock_result_set + + query_opts = {"optimizer_version": "latest"} + driver = SpannerSyncDriver( + connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} + ) + + driver.execute("SELECT 1") + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts + + +def test_config_provide_session_query_options() -> None: + """Verify config.provide_session forwards query_options to driver_features.""" + config = SpannerSyncConfig(connection_config={"project": "p", "instance_id": "i", "database_id": "d"}) + query_opts = {"optimizer_version": "5"} + features = config._session_driver_features( + request_options=None, directed_read_options=None, query_options=query_opts, retry=None, timeout=None + ) + assert features["query_options"] == query_opts + + +def test_driver_execute_many_omits_query_options() -> None: + """Verify execute_many does not forward query_options to batch_update.""" + mock_cursor = MagicMock() + mock_cursor.batch_update.return_value = (None, [1, 1]) + + query_opts = {"optimizer_version": "latest"} + driver = SpannerSyncDriver( + connection=mock_cursor, statement_config=default_statement_config, driver_features={"query_options": query_opts} + ) + + driver.execute_many("INSERT INTO t (id) VALUES (:id)", [{"id": 1}, {"id": 2}], query_options=query_opts) + + mock_cursor.batch_update.assert_called_once() + _, kwargs = mock_cursor.batch_update.call_args + assert "query_options" not in kwargs + + +def test_driver_select_stream_query_options() -> None: + """Verify select_stream passes query_options to execute_sql.""" + mock_cursor = MagicMock() + mock_result_set = MagicMock() + mock_field = MagicMock() + mock_field.name = "v" + mock_field.type_.code = 1 + mock_result_set.metadata.row_type.fields = [mock_field] + mock_result_set.__iter__.return_value = [[1], [2]] + mock_cursor.execute_sql.return_value = mock_result_set + + query_opts = {"optimizer_version": "6"} + driver = SpannerSyncDriver(connection=mock_cursor, statement_config=default_statement_config, driver_features={}) + + stream = driver.select_stream("SELECT 1", query_options=query_opts) + rows = list(stream) if stream else [] + assert len(rows) == 2 + + mock_cursor.execute_sql.assert_called_once() + _, kwargs = mock_cursor.execute_sql.call_args + assert kwargs.get("query_options") == query_opts + + +def test_driver_dml_retains_query_options() -> None: + cursor = MagicMock() + cursor.execute_update.return_value = 1 + driver = SpannerSyncDriver(connection=cursor) + driver.execute("UPDATE t SET x = 1", query_options={"optimizer_version": "6"}) + assert cursor.execute_update.call_args.kwargs["query_options"] == {"optimizer_version": "6"} From f5a35fd97fa2e6971a829d9238ba9e457b6aadfb Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:05:46 +0000 Subject: [PATCH 12/15] docs: reconcile unreleased adapter changelog after rebase --- docs/changelog.rst | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 43137c029..eac7e3f0c 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -55,16 +55,19 @@ Unreleased **Fixed:** -* SQL Server event queue DDL guards use the configured table and index names. - Arrow ODBC index checks no longer include column text in the table name. - -* Arrow read failures in the mssql-python adapter use SQLSpec error types. +* Spanner binds Decimal values as NUMERIC and boolean arrays as BOOL. + Typed null dictionaries use JSON, and JSON null results stay ``None``. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. * Arrow ODBC pagination reuses compiled placeholder positions instead of parsing SQL again. ADBC keeps bound values in its ADK store queries. DuckDB Arrow loads keep sparse dictionary fields and quote table names. + +* SQL Server event queue DDL guards use the configured table and index names. + Arrow ODBC index checks no longer include column text in the table name. + +* Arrow read failures in the mssql-python adapter use SQLSpec error types. * SQLite pools replace lost in-memory connections. Arrow imports roll back writes on failure or cancellation when the adapter owns the transaction. From adaaed6944f105c58a13ddb0ce1760b0a11078ef Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 23:37:02 +0000 Subject: [PATCH 13/15] test: register Spanner native query options feature --- tests/integration/adapters/_shared/_driver_type_system.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 540f43c86..c335d89ef 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -243,6 +243,7 @@ class SourceEquivalenceCase: "retry", "timeout", "request_options", + "query_options", "directed_read_options", "session_labels", "enable_events", From 553e62421112553f7aa3e3fbcd363ff55a8c4adf Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 28 Sep 2026 22:37:46 +0000 Subject: [PATCH 14/15] fix: resolve post-merge release review findings and align ADK 2.9/2.10 contracts --- docs/changelog.rst | 118 +++-- docs/reference/adapters/arrow_odbc.rst | 4 +- pyproject.toml | 1 + sqlspec/adapters/adbc/_typing.py | 3 +- sqlspec/adapters/adbc/adk/store.py | 78 +-- sqlspec/adapters/adbc/core.py | 16 +- sqlspec/adapters/aiomysql/adk/store.py | 11 +- sqlspec/adapters/aiomysql/core.py | 12 +- sqlspec/adapters/aiosqlite/_typing.py | 3 +- sqlspec/adapters/aiosqlite/adk/store.py | 90 ++-- sqlspec/adapters/aiosqlite/driver.py | 4 + sqlspec/adapters/arrow_odbc/_typing.py | 3 +- sqlspec/adapters/arrow_odbc/adk/store.py | 30 +- sqlspec/adapters/arrow_odbc/driver.py | 13 +- sqlspec/adapters/asyncmy/adk/store.py | 11 +- sqlspec/adapters/asyncmy/core.py | 13 +- sqlspec/adapters/asyncmy/driver.py | 2 +- sqlspec/adapters/asyncpg/_typing.py | 8 +- sqlspec/adapters/asyncpg/config.py | 17 +- sqlspec/adapters/bigquery/_typing.py | 22 +- sqlspec/adapters/bigquery/config.py | 7 +- sqlspec/adapters/bigquery/core.py | 31 +- sqlspec/adapters/bigquery/driver.py | 26 +- sqlspec/adapters/bigquery/litestar/store.py | 24 +- sqlspec/adapters/cockroach_asyncpg/_typing.py | 8 +- sqlspec/adapters/cockroach_psycopg/_typing.py | 8 +- .../adapters/cockroach_psycopg/adk/store.py | 13 +- sqlspec/adapters/cockroach_psycopg/driver.py | 21 +- sqlspec/adapters/db2/_typing.py | 12 +- sqlspec/adapters/db2/driver.py | 14 +- sqlspec/adapters/db2/pool.py | 10 +- sqlspec/adapters/duckdb/_typing.py | 3 +- sqlspec/adapters/duckdb/adk/store.py | 45 +- sqlspec/adapters/duckdb/driver.py | 43 +- sqlspec/adapters/duckdb/litestar/store.py | 6 +- sqlspec/adapters/mssql_python/adk/store.py | 41 +- sqlspec/adapters/mssql_python/config.py | 1 + sqlspec/adapters/mssql_python/core.py | 4 +- sqlspec/adapters/mssql_python/driver.py | 5 +- sqlspec/adapters/mysqlconnector/_typing.py | 4 +- sqlspec/adapters/mysqlconnector/adk/store.py | 22 +- sqlspec/adapters/mysqlconnector/config.py | 5 +- sqlspec/adapters/mysqlconnector/core.py | 24 +- sqlspec/adapters/oracledb/_json_handlers.py | 31 +- sqlspec/adapters/oracledb/adk/store.py | 450 +++++++++--------- sqlspec/adapters/oracledb/core.py | 122 ++--- sqlspec/adapters/oracledb/litestar/store.py | 144 +++--- sqlspec/adapters/oracledb/type_converter.py | 17 +- sqlspec/adapters/psqlpy/_typing.py | 3 +- sqlspec/adapters/psqlpy/adk/store.py | 40 +- sqlspec/adapters/psqlpy/config.py | 4 +- sqlspec/adapters/psqlpy/core.py | 11 +- sqlspec/adapters/psycopg/_typing.py | 3 +- sqlspec/adapters/psycopg/adk/store.py | 78 ++- sqlspec/adapters/psycopg/config.py | 12 +- sqlspec/adapters/psycopg/type_converter.py | 8 +- sqlspec/adapters/pymssql/adk/store.py | 17 +- sqlspec/adapters/pymssql/core.py | 2 +- sqlspec/adapters/pymssql/driver.py | 5 +- sqlspec/adapters/pymssql/pool.py | 8 +- sqlspec/adapters/pymysql/_typing.py | 37 +- sqlspec/adapters/pymysql/adk/store.py | 11 +- sqlspec/adapters/pymysql/config.py | 7 +- sqlspec/adapters/pymysql/core.py | 12 +- sqlspec/adapters/pymysql/litestar/store.py | 7 +- sqlspec/adapters/spanner/adk/store.py | 47 +- sqlspec/adapters/spanner/config.py | 13 +- sqlspec/adapters/spanner/core.py | 7 - sqlspec/adapters/spanner/driver.py | 9 +- sqlspec/adapters/spanner/type_converter.py | 6 +- sqlspec/adapters/sqlite/_typing.py | 3 +- sqlspec/adapters/sqlite/adk/store.py | 34 +- sqlspec/adapters/sqlite/driver.py | 4 + sqlspec/builder/_base.py | 21 +- sqlspec/builder/_merge.py | 10 +- sqlspec/dialects/spanner/_parsers.py | 2 +- sqlspec/extensions/adk/memory/converters.py | 16 +- sqlspec/extensions/adk/memory/service.py | 102 +++- sqlspec/extensions/adk/service.py | 20 +- sqlspec/migrations/base.py | 15 +- sqlspec/migrations/schema.py | 28 +- sqlspec/migrations/tracker.py | 5 +- .../adapters/test_aiomysql/test_adk_store.py | 100 ++++ .../adapters/test_asyncmy/test_adk_store.py | 101 ++++ .../unit/adapters/test_asyncmy/test_driver.py | 25 +- .../adapters/test_bigquery/test_adk_store.py | 66 +++ .../unit/adapters/test_bigquery/test_core.py | 17 + .../test_bigquery/test_storage_write_api.py | 8 +- tests/unit/adapters/test_db2/test_pool.py | 2 + .../adapters/test_db2/test_transactions.py | 25 + .../test_mssql_python/test_adk_store.py | 49 +- .../adapters/test_mssql_python/test_config.py | 13 + .../adapters/test_mssql_python/test_core.py | 10 +- .../test_transaction_state.py | 26 +- .../test_mysqlconnector/test_adk_store.py | 150 ++++++ .../test_mysqlconnector/test_config.py | 18 + .../adapters/test_mysqlconnector/test_core.py | 44 +- .../test_oracledb/test_litestar_store.py | 12 + .../test_oracledb/test_lob_coercion.py | 97 +++- .../test_oracledb/test_oracle_adk_store.py | 34 ++ tests/unit/adapters/test_psqlpy/test_core.py | 18 +- .../adapters/test_psycopg/test_adk_store.py | 50 ++ .../adapters/test_pymssql/test_adk_store.py | 37 ++ tests/unit/adapters/test_pymssql/test_core.py | 6 + .../unit/adapters/test_pymssql/test_driver.py | 2 +- .../adapters/test_pymssql/test_extensions.py | 36 +- tests/unit/adapters/test_pymssql/test_pool.py | 2 + .../adapters/test_pymysql/test_adk_store.py | 90 ++++ .../test_pymysql/test_cloud_sql_connector.py | 4 +- .../adapters/test_spanner/test_adk_store.py | 14 + .../test_spanner/test_batch_write_api.py | 2 +- .../unit/adapters/test_spanner/test_config.py | 3 + .../test_load_from_arrow_mutations.py | 31 +- tests/unit/builder/test_merge.py | 23 + tests/unit/dialects/test_spanner_hints.py | 13 + .../test_adk/test_memory_converters.py | 59 +++ .../unit/extensions/test_adk/test_service.py | 18 +- tests/unit/migrations/test_schema_ensure.py | 18 + tools/scripts/mypyc_inventory.py | 120 +++-- tools/scripts/mypyc_smoke.py | 2 + 120 files changed, 2516 insertions(+), 1036 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index eac7e3f0c..a6bfdcff2 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -17,104 +17,162 @@ Unreleased * BigQuery supports native query resource controls, explicit STRUCT parameters, typed empty arrays, and configurable Storage Write stream modes while retaining the atomic PENDING default. + (`#812 `_) * The mssql-python adapter can load Arrow streams with native BulkCopy options. Columns map by name by default, and overwrite still uses DELETE. + (`#819 `_) * Pymssql connection types include native encryption settings. - + (`#819 `_) * SQLite and aiosqlite can register custom window functions on Python 3.11 and later when the SQLite runtime supports them. Choose a transaction lock mode or set the batch size for Arrow imports. Defaults stay the same. - -* Spanner forwards native query options and final-statement hints and - supports opt-in Batch Write from read sessions. - + (`#820 `_) +* Spanner forwards native query options (``QueryOptions``, ``RequestOptions``, + ``DirectedReadOptions``) and final-statement hints, supports opt-in Batch Write + from read sessions, adds ``run_in_transaction`` retry with backoff on ``Aborted`` + and ``last_statement=True`` commit inlining, supports multiplexed session pooling + and ``PingingPool`` keepalive intervals, coerces ``FLOAT32`` vector parameters, + and adds ``execute_partitioned_dml()`` for Partitioned DML truncation during + ``load_from_arrow(overwrite=True)``. + (`#814 `_) +* Cloud Spanner and Spangres SQLGlot dialects isolate custom ``SpannerParser``, + ``SpangresParser``, ``SpannerGenerator``, and ``SpangresGenerator`` subclasses + and expand AST and transpilation support for ``INTERLEAVE IN PARENT`` with + ``ON DELETE``, ``TTL``, ``ROW DELETION POLICY``, ``SEARCH`` / ``SCORE`` / + ``SNIPPETS`` / ``TOKENLIST``, ``VECTOR_INDEX`` with ``OPTIONS``, + ``GRAPH_TABLE``, ``FLOAT32``, ``SAFE_CAST`` / ``TRY_CAST``, + ``SPANNER.ML_PREDICT_ROW``, Spanner sequences, and ``spangres`` DDL/DML + transpilation and catalog query packs. + (`#813 `_) * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. - + (`#818 `_) * ADBC FlightSQL adds options for TLS/mTLS, RPC timeouts, message size, cookies and headers. Values set in native ``db_kwargs`` take precedence. - + (`#818 `_) * PostgreSQL adapters expose native asyncpg custom codecs and per-query timeouts, psycopg null pools and JSON codecs, supported CockroachDB startup settings, and psqlpy dense-vector conversion. PgBouncer compatibility mode avoids explicit prepared stack statements without weakening transaction cleanup. Null pools preserve concurrency limits, and timeout forwarding retains explicit zero values. - + (`#822 `_) * Added an IBM Db2 adapter for Db2 LUW 11.5 and later with sync (``Db2SyncConfig``) and async (``Db2AsyncConfig``) configurations built on ``ibm_db``. It includes connection pooling, catalog reflection, migrations, and Litestar session, events queue, and Google ADK stores. See :doc:`reference/adapters/db2`. + (`#811 `_) * Added a ``db2`` SQL dialect. It renders Db2 paging, special registers, labeled durations, isolation and lock clauses, and Db2 data types, and translates builder row locks to Db2 lock clauses. + (`#811 `_) * The arrow-odbc adapter supports IBM Db2 through the IBM CLI/ODBC driver, including Db2 connection keywords, transactions, lowercase result columns, and its Litestar session, events queue, and Google ADK stores. + (`#811 `_) + +**Changed:** + +* Extracted shared MySQL driver primitives into ``sqlspec/adapters/mysql_common.py`` + across ``aiomysql``, ``asyncmy``, ``pymysql``, and ``mysqlconnector``, and + reorganized dialect data-dictionary definitions into dialect-scoped + ``sqlspec/data_dictionary/dialects//`` packages with shared MySQL + data-dictionary base classes. + (`#821 `_) +* Simplified OracleDB and Db2 typing and helper boundaries: encapsulated + ``OraclePipelineDriver`` protocol typing, exported ``Db2ConnectionParams`` and + ``Db2DriverFeatures``, removed redundant ``if not TYPE_CHECKING:`` typing + blocks in ``oracledb`` and ``db2``, internalized adapter core helpers, and + made Oracle JSON, UUID, and vector type-handler registration idempotent. + (`#817 `_, + `#823 `_) **Fixed:** * Spanner binds Decimal values as NUMERIC and boolean arrays as BOOL. Typed null dictionaries use JSON, and JSON null results stay ``None``. + (`#814 `_) * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. - + (`#818 `_) * Arrow ODBC pagination reuses compiled placeholder positions instead of parsing SQL again. ADBC keeps bound values in its ADK store queries. DuckDB Arrow loads keep sparse dictionary fields and quote table names. - -* SQL Server event queue DDL guards use the configured table and index names. - Arrow ODBC index checks no longer include column text in the table name. - + (`#818 `_) +* T-SQL event store ``CREATE INDEX`` ``OBJECT_ID`` guards across ``arrow_odbc``, + ``pymssql``, and ``mssql_python`` use the configured table and generated index + names directly so ``arrow_odbc`` index checks no longer capture column list + text inside the table name. + (`#816 `_, + `#819 `_, + `#828 `_, + `#829 `_) * Arrow read failures in the mssql-python adapter use SQLSpec error types. + (`#819 `_) * SQLite pools replace lost in-memory connections. Arrow imports roll back writes on failure or cancellation when the adapter owns the transaction. - + (`#820 `_) * Builder results keep CTE trees independent, and column pruning no longer exposes its cached expression to mutation. SQL generation avoids redundant copies of temporary trees while preserving caller and cache ownership. - + (`#829 `_) * SQL Server migration drivers retain the previous default schema if restoring it fails, so cleanup can be retried. The migration guide clarifies that this setting belongs to the database user rather than one connection. - + (`#829 `_) * Asyncpg stack telemetry reports sequential prepared execution rather than native pipelining. Each statement still returns its own result. - -* Fixture files keep JSON strings such as ``"true"`` and ``"[1]"`` as strings. - This also works for SQLite JSON columns. Write objects and arrays directly - instead of encoding them as strings. Column filtering respects case. - See :doc:`usage/testing`. - -* DDL builders and migration trackers keep quoted table names intact. - Names with spaces and mixed-case Oracle names retain their quotes. - + (`#822 `_, + `#829 `_) +* Fixture files keep JSON strings such as ``"true"`` and ``"[1]"`` as strings + during export and load round-trips across DuckDB, SQLite, ADBC, and MySQL + (``#824``), normalize bare string ``conflict_keys`` values and ignore tables + absent from subset loads (``#825``), and support sparse row keys, + ``ignore_unknown_columns``, and ``exclude_update_columns`` for upserts + (``#826``). Column filtering respects case. See :doc:`usage/testing`. + (`#827 `_, + `#829 `_) +* Quoted and schema-qualified ``version_table`` identifiers are preserved across + migration trackers and DDL builders (``CreateTable``, ``DropTable``, + ``AlterTable``). Trackers keep unquoted catalog lookup names in + ``version_table_name`` and ``version_table_schema`` while retaining exact + identifier quotes in DDL and tracking queries, including names with spaces and + mixed-case Oracle identifiers. + (`#815 `_, + `#828 `_, + `#829 `_) * MySQL pools release connections when setup fails. They discard connections that fail to roll back. MySQL Connector keeps native async pooling on Connector 9.4 and later, plus direct connections on older versions. Asyncmy retains native ``LOAD DATA LOCAL INFILE`` support. - + (`#821 `_, + `#829 `_) * Oracle keeps Thick-mode options for sync pools. Async pools reject Thick mode before they open. Pool shutdown preserves native checks for borrowed connections. Custom handlers still convert LOBs, and JSON handlers preserve the user's callbacks. - + (`#817 `_, + `#829 `_) * Db2 pools clean up after failed or cancelled setup. Batch results keep an unknown row count when the driver cannot report one. String searches keep their start position and requested occurrence. - + (`#811 `_, + `#829 `_) * Spanner schema queries no longer require a table name. SQL output keeps JOIN hints and plain comments. Sequence statements keep qualified names and ``IF NOT EXISTS`` guards. Cached row converters refresh when the configured JSON deserializer changes. - + (`#813 `_, + `#829 `_) * Psycopg reads COPY files in chunks, not all at once. ADK stores use RETURNING to cut round trips. Psqlpy closes a connection if setup fails. - + (`#822 `_) * Builder upserts emit ``MERGE`` for the ``db2`` dialect. + (`#811 `_) * The arrow-odbc adapter detects the SQL dialect from the ODBC driver name only, so database, host, or user names no longer select the wrong dialect. + (`#811 `_) v0.64.0 - Startup performance, connection normalization, and adapter lifecycle hardening ----------------------------------------------------------------------------------------- diff --git a/docs/reference/adapters/arrow_odbc.rst b/docs/reference/adapters/arrow_odbc.rst index 08403f15f..c7bbe2e87 100644 --- a/docs/reference/adapters/arrow_odbc.rst +++ b/docs/reference/adapters/arrow_odbc.rst @@ -10,8 +10,8 @@ PostgreSQL, MySQL, or other ODBC sources and the Arrow ecosystem. SQL Server coverage is exercised in CI against SQL Server 2022 through ``pytest-databases`` and Microsoft ODBC Driver 18. The shared contract matrix verifies native Arrow reads, Arrow reader/batch output, and Arrow bulk ingest -for this adapter. Row-oriented ``execute_many()`` is intentionally unsupported; -use ``load_from_arrow()`` for bulk writes. +for this adapter. Use ``load_from_arrow()`` or ``bulk_insert_arrow()`` for +native Arrow bulk writes, or ``execute_many()`` for row-oriented batches. The adapter exports a table-backed events queue store, a Litestar session store, and Google ADK session/event and memory stores. They support SQL Server diff --git a/pyproject.toml b/pyproject.toml index 6158692aa..c6f6ec0c6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -236,6 +236,7 @@ exclude = [ "sqlspec/dialects/postgres/_pgvector.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/postgres/_paradedb.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/postgres/_pg_textsearch.py", # interpreted: Dialect subclass needs sqlglot's metaclass + "sqlspec/dialects/spanner/_expressions.py", # interpreted: SQLGlot AST expression helpers stay interpreted "sqlspec/dialects/spanner/_spanner.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/dialects/spanner/_spangres.py", # interpreted: Dialect subclass needs sqlglot's metaclass "sqlspec/storage/_arrow_payload.py", # PyArrow conversion boundary stays interpreted diff --git a/sqlspec/adapters/adbc/_typing.py b/sqlspec/adapters/adbc/_typing.py index 292a6b68b..ccdcc896d 100644 --- a/sqlspec/adapters/adbc/_typing.py +++ b/sqlspec/adapters/adbc/_typing.py @@ -31,8 +31,7 @@ AdbcConnection: TypeAlias = _AdbcConnection AdbcRawCursor: TypeAlias = _AdbcRawCursor AdbcNativeError: TypeAlias = _AdbcNativeError - -if not TYPE_CHECKING: +else: AdbcConnection = _AdbcConnection AdbcRawCursor = _AdbcRawCursor AdbcNativeError = _AdbcNativeError diff --git a/sqlspec/adapters/adbc/adk/store.py b/sqlspec/adapters/adbc/adk/store.py index 42fb5b783..84d8bd101 100644 --- a/sqlspec/adapters/adbc/adk/store.py +++ b/sqlspec/adapters/adbc/adk/store.py @@ -988,37 +988,41 @@ def _append_event_and_update_state( with self._config.provide_connection() as conn: cursor = conn.cursor() try: - self._execute( - cursor, - insert_sql, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - event_data, - ), - ) self._execute(cursor, update_sql, (state_json, app_name, user_id, session_id)) - if app_state is not None: - self._execute(cursor, delete_app_state_sql, (app_name,)) - self._execute( - cursor, - insert_app_state_sql, - (app_name, self._serialize_state(app_state), datetime.now(timezone.utc)), - ) - if user_state is not None: - self._execute(cursor, delete_user_state_sql, (app_name, user_id)) + self._execute(cursor, select_sql, (app_name, user_id, session_id)) + row = cursor.fetchone() + if row is not None: self._execute( cursor, - insert_user_state_sql, - (app_name, user_id, self._serialize_state(user_state), datetime.now(timezone.utc)), + insert_sql, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + event_data, + ), ) - self._execute(cursor, select_sql, (app_name, user_id, session_id)) - row = cursor.fetchone() - conn.commit() + if app_state is not None: + self._execute(cursor, delete_app_state_sql, (app_name,)) + self._execute( + cursor, + insert_app_state_sql, + (app_name, self._serialize_state(app_state), datetime.now(timezone.utc)), + ) + if user_state is not None: + self._execute(cursor, delete_user_state_sql, (app_name, user_id)) + self._execute( + cursor, + insert_user_state_sql, + (app_name, user_id, self._serialize_state(user_state), datetime.now(timezone.utc)), + ) + conn.commit() + else: + with contextlib.suppress(Exception): + conn.rollback() except Exception: with contextlib.suppress(Exception): conn.rollback() @@ -1304,6 +1308,9 @@ def dialect(self) -> str: def create_tables(self) -> None: """Create tables if they don't exist.""" + if not self._enabled: + return + if not self.create_schema_enabled: self.reconcile_schema() return @@ -1583,7 +1590,7 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec try: for entry in entries: content_json = self._serialize_json_field(entry["content_json"]) - metadata_json = self._serialize_json_field(entry["metadata_json"]) + metadata_json = self._serialize_json_field(entry.get("metadata_json")) params: tuple[Any, ...] if self._owner_id_column_name: params = ( @@ -1593,7 +1600,7 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, self._encode_timestamp(entry["timestamp"]), content_json, @@ -1609,7 +1616,7 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), self._encode_timestamp(entry["timestamp"]), content_json, entry["content_text"], @@ -1647,6 +1654,9 @@ def _search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) + if not query or not query.strip(): + return [] + if self._use_fts: logger.warning("ADBC memory store does not support FTS, falling back to simple search") @@ -1681,6 +1691,10 @@ def _search_entries( return self._rows_to_records(rows) def _delete_entries_by_session(self, session_id: str) -> int: + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + use_returning = self._dialect in {DIALECT_SQLITE, DIALECT_POSTGRESQL, DIALECT_DUCKDB} if use_returning: sql = f"DELETE FROM {self._memory_table} WHERE session_id = ? RETURNING 1" @@ -1700,6 +1714,10 @@ def _delete_entries_by_session(self, session_id: str) -> int: cursor.close() def _delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + cutoff = self._encode_timestamp(datetime.now(timezone.utc) - timedelta(days=days)) use_returning = self._dialect in {DIALECT_SQLITE, DIALECT_POSTGRESQL, DIALECT_DUCKDB} clauses = ["inserted_at < ?"] diff --git a/sqlspec/adapters/adbc/core.py b/sqlspec/adapters/adbc/core.py index 713d9f29b..1094f8121 100644 --- a/sqlspec/adapters/adbc/core.py +++ b/sqlspec/adapters/adbc/core.py @@ -11,6 +11,7 @@ from sqlglot import exp from sqlglot.errors import ParseError +import sqlspec.adapters.adbc._typing as _adbc_typing from sqlspec.adapters.adbc.type_converter import get_adbc_type_converter from sqlspec.core import ( DriverParameterProfile, @@ -693,7 +694,6 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec sqlstate = sqlstate_attr if sqlstate_attr is not None else None if sqlstate: - # Use centralized SQLSTATE mapping for specific codes if sqlstate == "23505": return _create_adbc_error(error, UniqueViolationError, "unique constraint violation") if sqlstate == "23503": @@ -703,7 +703,6 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec if sqlstate == "23514": return _create_adbc_error(error, CheckViolationError, "check constraint violation") - # Deadlock and serialization errors if sqlstate == "40P01": return _create_adbc_error(error, DeadlockError, "deadlock detected") if sqlstate == "40001": @@ -713,25 +712,20 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec termination_class = _classify_timeout_or_cancellation(str(error)) or OperationalError return _create_adbc_error(error, termination_class, "query terminated") - # Permission errors if sqlstate == "42501": return _create_adbc_error(error, PermissionDeniedError, "insufficient privilege") if sqlstate == "28000": return _create_adbc_error(error, PermissionDeniedError, "invalid authorization") - # Use centralized mapping for SQLSTATE class prefixes exc_class = map_sqlstate_to_exception(sqlstate) if exc_class is not None and exc_class is not SQLSpecError: description = _get_sqlstate_description(sqlstate) return _create_adbc_error(error, exc_class, description) - # Fallback for unmapped SQLSTATE codes return _create_adbc_error(error, SQLSpecError, "database error") - # Message-based fallback when no SQLSTATE is available error_msg = str(error).lower() - # Constraint violations if "unique" in error_msg or "duplicate" in error_msg: return _create_adbc_error(error, UniqueViolationError, "unique constraint violation") if "foreign key" in error_msg: @@ -743,7 +737,6 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec if "constraint" in error_msg: return _create_adbc_error(error, IntegrityError, "integrity constraint violation") - # Deadlock/lock patterns if "deadlock" in error_msg: return _create_adbc_error(error, DeadlockError, "deadlock detected") if "serialization" in error_msg or "concurrent update" in error_msg: @@ -752,15 +745,12 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec if message_class := _classify_timeout_or_cancellation(error_msg): return _create_adbc_error(error, message_class, "query terminated") - # Permission patterns if "permission" in error_msg or "denied" in error_msg or "unauthorized" in error_msg: return _create_adbc_error(error, PermissionDeniedError, "permission denied") - # Syntax errors if "syntax" in error_msg: return _create_adbc_error(error, SQLParsingError, "SQL parsing error") - # Connection errors if "connection" in error_msg or "connect" in error_msg: return _create_adbc_error(error, DatabaseConnectionError, "connection error") @@ -1396,10 +1386,10 @@ def _prepare_batch_with_casts( def _resolve_flightsql_db_kwargs_keys() -> "tuple[str, str]": try: - from sqlspec.adapters.adbc._typing import AdbcFlightSqlDatabaseOptions as DatabaseOptions + database_options = _adbc_typing.AdbcFlightSqlDatabaseOptions except ImportError: return _FLIGHTSQL_TLS_SKIP_VERIFY_KEY, _FLIGHTSQL_AUTHORIZATION_HEADER_KEY - return DatabaseOptions.TLS_SKIP_VERIFY.value, DatabaseOptions.AUTHORIZATION_HEADER.value + return database_options.TLS_SKIP_VERIFY.value, database_options.AUTHORIZATION_HEADER.value def _lift_flightsql_db_kwargs(config: "dict[str, Any]") -> None: diff --git a/sqlspec/adapters/aiomysql/adk/store.py b/sqlspec/adapters/aiomysql/adk/store.py index f3ac15c49..6e943d65d 100644 --- a/sqlspec/adapters/aiomysql/adk/store.py +++ b/sqlspec/adapters/aiomysql/adk/store.py @@ -579,12 +579,12 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) else: @@ -595,11 +595,11 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) await cursor.execute(sql, params) @@ -655,6 +655,9 @@ async def search_entries( records: list[StoredMemory] = [] for row in rows: rec = cast("StoredMemory", dict(zip(columns, row, strict=False))) + rec["content_json"] = _json_dict(rec.get("content_json")) + metadata_val = rec.get("metadata_json") + rec["metadata_json"] = _json_dict(metadata_val) if metadata_val is not None else None rec["embedding"] = None records.append(rec) return records diff --git a/sqlspec/adapters/aiomysql/core.py b/sqlspec/adapters/aiomysql/core.py index ecc834bb2..be8af2da1 100644 --- a/sqlspec/adapters/aiomysql/core.py +++ b/sqlspec/adapters/aiomysql/core.py @@ -84,10 +84,14 @@ async def start(self) -> None: self._driver._check_pending_exception(handler) async def _start(self) -> None: - cursor = await self._driver.connection.cursor(AiomysqlSSCursor) - self._cursor = cursor - await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + try: + cursor = await self._driver.connection.cursor(AiomysqlSSCursor) + self._cursor = cursor + await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) + self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + except BaseException: + await self.close(error=True) + raise async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() diff --git a/sqlspec/adapters/aiosqlite/_typing.py b/sqlspec/adapters/aiosqlite/_typing.py index 55b6f9588..8d0168ba9 100644 --- a/sqlspec/adapters/aiosqlite/_typing.py +++ b/sqlspec/adapters/aiosqlite/_typing.py @@ -27,8 +27,7 @@ AiosqliteConnection: TypeAlias = _AiosqliteConnection AiosqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] AiosqliteRawCursor: TypeAlias = aiosqlite.Cursor - -if not TYPE_CHECKING: +else: AiosqliteConnection = _AiosqliteConnection AiosqliteConnectionFactory = TypeAliasType("AiosqliteConnectionFactory", type[sqlite3.Connection]) AiosqliteRawCursor = aiosqlite.Cursor diff --git a/sqlspec/adapters/aiosqlite/adk/store.py b/sqlspec/adapters/aiosqlite/adk/store.py index f863a2c53..df7990113 100644 --- a/sqlspec/adapters/aiosqlite/adk/store.py +++ b/sqlspec/adapters/aiosqlite/adk/store.py @@ -13,9 +13,11 @@ from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseAsyncADKMemoryStore +from sqlspec.utils.logging import get_logger from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: + import logging from collections.abc import Sequence from sqlspec.adapters.aiosqlite.config import AiosqliteConfig @@ -30,6 +32,8 @@ _FTS_DETAIL_VALUES: Final = frozenset({"full", "column", "none"}) _FTS_TOKENIZE_PATTERN: Final = re.compile(r"^[A-Za-z0-9_ -]+$") +logger: "logging.Logger" = get_logger("sqlspec.adapters.aiosqlite.adk.store") + class AiosqliteADKConfig(ADKConfig): """Aiosqlite-specific ADK extension settings. @@ -776,11 +780,11 @@ async def create_tables(self) -> None: Skips table creation if memory store is disabled. """ - if not self.create_schema_enabled: - await self.reconcile_schema() + if not self._enabled: return - if not self._enabled: + if not self.create_schema_enabled: + await self.reconcile_schema() return async with self._config.provide_session() as driver: @@ -811,6 +815,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " """ for entry in entries: scope = entry.get("scope", "user") + metadata_json = entry.get("metadata_json") params_list.append(( entry["id"], entry["session_id"], @@ -818,12 +823,12 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], scope, entry["event_id"], - entry["author"], + entry.get("author"), owner_id, _datetime_to_julian(entry["timestamp"]), to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(metadata_json) if metadata_json is not None else None, _datetime_to_julian(entry["inserted_at"]), )) else: @@ -835,6 +840,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " """ for entry in entries: scope = entry.get("scope", "user") + metadata_json = entry.get("metadata_json") params_list.append(( entry["id"], entry["session_id"], @@ -842,11 +848,11 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], scope, entry["event_id"], - entry["author"], + entry.get("author"), _datetime_to_julian(entry["timestamp"]), to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(metadata_json) if metadata_json is not None else None, _datetime_to_julian(entry["inserted_at"]), )) cursor = await conn.executemany(sql, params_list) @@ -871,35 +877,55 @@ async def search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) - if not query: + if not query or not query.strip(): return [] - limit_value = limit or self._max_results - if self._use_fts: - where_scope, scope_params = _build_sqlite_scope_clause("m.", app_name, user_id, scope_filter) - sql = f""" - SELECT m.* FROM {self._memory_table} AS m - JOIN {self._memory_table}_fts AS fts ON m.rowid = fts.rowid - WHERE {where_scope} AND fts.content_text MATCH ? - ORDER BY m.timestamp DESC - LIMIT ? - """ - params = (*scope_params, query, limit_value) - else: - where_scope, scope_params = _build_sqlite_scope_clause("", app_name, user_id, scope_filter) - sql = f""" - SELECT * FROM {self._memory_table} - WHERE {where_scope} AND content_text LIKE ? - ORDER BY timestamp DESC - LIMIT ? - """ - params = (*scope_params, f"%{query}%", limit_value) + limit_value = limit if limit is not None else self._max_results + rows: list[Any] | None = None + columns: list[str] = [] async with self._config.provide_connection() as conn: - cursor = await conn.execute(sql, params) - rows = await cursor.fetchall() - columns = [col[0] for col in cursor.description or []] - await cursor.close() + if self._use_fts: + where_scope, scope_params = _build_sqlite_scope_clause("m.", app_name, user_id, scope_filter) + fts_sql = f""" + SELECT m.* FROM {self._memory_table} AS m + JOIN {self._memory_table}_fts AS fts ON m.rowid = fts.rowid + WHERE {where_scope} AND fts.content_text MATCH ? + ORDER BY m.timestamp DESC + LIMIT ? + """ + fts_params = (*scope_params, query, limit_value) + try: + cursor = await conn.execute(fts_sql, fts_params) + try: + rows = list(await cursor.fetchall()) + columns = [col[0] for col in cursor.description or []] + finally: + await cursor.close() + except Exception as exc: + logger.warning("FTS search failed; falling back to simple search: %s", exc) + + if rows is None: + where_scope, scope_params = _build_sqlite_scope_clause("", app_name, user_id, scope_filter) + sql = f""" + SELECT * FROM {self._memory_table} + WHERE {where_scope} AND content_text LIKE ? + ORDER BY timestamp DESC + LIMIT ? + """ + params = (*scope_params, f"%{query}%", limit_value) + try: + cursor = await conn.execute(sql, params) + try: + rows = list(await cursor.fetchall()) + columns = [col[0] for col in cursor.description or []] + finally: + await cursor.close() + except sqlite3.OperationalError as exc: + if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): + return [] + raise + records: list[StoredMemory] = [] for row in rows: raw = dict(zip(columns, row, strict=False)) diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index b7ffe5f66..7d8b03d17 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -457,6 +457,10 @@ def _can_use_execute_many_thin_path( return False if "?" not in statement: return False + if self._resolve_dml_operation_type(statement) not in {"INSERT", "UPDATE", "DELETE"}: + return False + if "RETURNING" in statement.upper(): + return False parameter_config = config.parameter_config if parameter_config.default_parameter_style is not ParameterStyle.QMARK: diff --git a/sqlspec/adapters/arrow_odbc/_typing.py b/sqlspec/adapters/arrow_odbc/_typing.py index 8be6f9a2b..0be32aa33 100644 --- a/sqlspec/adapters/arrow_odbc/_typing.py +++ b/sqlspec/adapters/arrow_odbc/_typing.py @@ -21,8 +21,7 @@ ArrowOdbcConnection: TypeAlias = _arrow_odbc.Connection ArrowOdbcRawCursor: TypeAlias = _arrow_odbc.Connection - -if not TYPE_CHECKING: +else: ArrowOdbcConnection = _arrow_odbc.Connection ArrowOdbcRawCursor = _arrow_odbc.Connection diff --git a/sqlspec/adapters/arrow_odbc/adk/store.py b/sqlspec/adapters/arrow_odbc/adk/store.py index 26f1e2da8..c15674f46 100644 --- a/sqlspec/adapters/arrow_odbc/adk/store.py +++ b/sqlspec/adapters/arrow_odbc/adk/store.py @@ -449,12 +449,13 @@ def __init__(self, config: "ArrowOdbcConfig") -> None: def create_tables(self) -> None: """Create the memory table and indexes the catalog reports as missing.""" + if not self._enabled: + return + if not self.create_schema_enabled: self.reconcile_schema() return - if not self._enabled: - return with self._config.provide_session() as driver: self._sql.create_missing_objects( driver, ((self._memory_table, self._memory_table_ddl()),), self._memory_index_specs() @@ -516,17 +517,28 @@ def search_entries( if not self._enabled: msg = "Memory store is disabled" raise RuntimeError(msg) + if not query or not query.strip(): + return [] effective_limit = max(0, int(limit if limit is not None else self._max_results)) if effective_limit == 0: return [] where_scope, scope_params = _build_arrow_odbc_scope_where(app_name, user_id, scope_filter) - rows = self._execute_fetchall( - self._sql.memory_search_sql(self._memory_table, where_scope, effective_limit), (*scope_params, f"%{query}%") - ) + try: + rows = self._execute_fetchall( + self._sql.memory_search_sql(self._memory_table, where_scope, effective_limit), + (*scope_params, f"%{query}%"), + ) + except SQLSpecError as exc: + if self._sql.is_table_missing(exc): + return [] + raise return [_memory_record_from_row(row) for row in rows] def delete_entries_by_session(self, session_id: str) -> int: """Delete all memory entries for a specific session.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) table_ref = self._sql.table_ref(self._memory_table) count = self._select_count(f"SELECT COUNT(*) AS row_count FROM {table_ref} WHERE session_id = ?", (session_id,)) self._execute(f"DELETE FROM {table_ref} WHERE session_id = ?", (session_id,), commit=True) @@ -534,6 +546,9 @@ def delete_entries_by_session(self, session_id: str) -> int: def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: """Delete memory entries older than ``days`` days.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) cutoff = datetime.now(timezone.utc).timestamp() - (days * 86_400) cutoff_dt = datetime.fromtimestamp(cutoff, tz=timezone.utc) clauses = ["inserted_at < ?"] @@ -817,6 +832,7 @@ def _event_record_from_row(row: Any) -> StoredEvent: def _memory_insert_params( entry: StoredMemory, format_datetime: "Callable[[datetime | None], str | None]" ) -> "tuple[Any, ...]": + metadata_json = entry.get("metadata_json") return ( entry["id"], entry["session_id"], @@ -824,11 +840,11 @@ def _memory_insert_params( entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), format_datetime(entry["timestamp"]), to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]) if entry["metadata_json"] is not None else None, + to_json(metadata_json) if metadata_json is not None else None, format_datetime(entry["inserted_at"]), ) diff --git a/sqlspec/adapters/arrow_odbc/driver.py b/sqlspec/adapters/arrow_odbc/driver.py index ed40fc35b..fbf5748b9 100644 --- a/sqlspec/adapters/arrow_odbc/driver.py +++ b/sqlspec/adapters/arrow_odbc/driver.py @@ -28,15 +28,20 @@ from sqlspec.core.parameters._validator import ParameterValidator from sqlspec.driver import BaseSyncExceptionHandler, SyncDriverAdapterBase, SyncRowStream, validate_savepoint_name from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError +from sqlspec.typing import import_optional from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.text import quote_identifier, split_qualified_identifier if TYPE_CHECKING: + import pyarrow as pa + from sqlspec.builder import QueryBuilder from sqlspec.core import ArrowResult, ParameterProfile, Statement, StatementConfig, StatementFilter from sqlspec.driver import ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import ArrowRecordBatch, ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters +else: + pa = import_optional("pyarrow") __all__ = ("ArrowOdbcCursor", "ArrowOdbcDriver", "ArrowOdbcExceptionHandler", "resolve_dialect_from_dbms_name") @@ -194,7 +199,6 @@ def dispatch_execute_many(self, cursor: "ArrowOdbcRawCursor", statement: "SQL") query=sql, parameters=_odbc_parameters(parameter_set, naive_utc_datetimes=self._dialect == "db2") ) executed = True - # The native execute API supplies no affected-row count. return self.create_execution_result(cursor, rowcount_override=-1 if executed else 0, is_many_result=True) def dispatch_execute_script(self, cursor: "ArrowOdbcRawCursor", statement: "SQL") -> "ExecutionResult": @@ -375,7 +379,6 @@ def select_to_arrow( def bulk_insert_arrow(self, target_table: str, source: Any, *, chunk_size: int | None = None) -> None: """Insert an Arrow table or reader into a database table.""" ensure_pyarrow() - import pyarrow as pa resolved_chunk_size = chunk_size or self._chunk_size() exc_handler = self.handle_database_exceptions() @@ -467,8 +470,6 @@ def _normalize_table(self, table: Any) -> Any: def _normalize_reader(self, reader: "ArrowRecordBatchReader") -> "ArrowRecordBatchReader": """Wrap a record batch reader so its schema and batches carry normalized column names.""" - import pyarrow as pa - schema = reader.schema names = normalize_column_names(schema.names, self._lowercase_column_names) if names == schema.names: @@ -556,7 +557,6 @@ def _inline_mssql_pagination_parameters( or statement_config.parameter_config.output_transformer is not None ) ): - # Output transformers run after profile construction and can shift offsets. parameter_info = tuple(ParameterValidator(cache_max_size=0).extract_parameters(sql)) positions = {parameter.position: parameter.ordinal for parameter in parameter_info} replacements: list[tuple[int, str]] = [] @@ -620,7 +620,6 @@ def _odbc_parameters(parameters: Any, *, naive_utc_datetimes: bool = False) -> " def _reader_to_table(reader: Any) -> Any: ensure_pyarrow() - import pyarrow as pa if isinstance(reader, pa.Table): return reader @@ -637,7 +636,6 @@ def _reader_to_table(reader: Any) -> Any: def _to_pyarrow_reader(reader: object) -> "ArrowRecordBatchReader": ensure_pyarrow() - import pyarrow as pa if isinstance(reader, pa.RecordBatchReader): return reader @@ -659,7 +657,6 @@ def _to_pyarrow_reader(reader: object) -> "ArrowRecordBatchReader": def _table_to_reader(table: Any, chunk_size: int) -> Any: ensure_pyarrow() - import pyarrow as pa return pa.RecordBatchReader.from_batches(table.schema, table.to_batches(max_chunksize=chunk_size)) diff --git a/sqlspec/adapters/asyncmy/adk/store.py b/sqlspec/adapters/asyncmy/adk/store.py index 9d743b904..8fe51925e 100644 --- a/sqlspec/adapters/asyncmy/adk/store.py +++ b/sqlspec/adapters/asyncmy/adk/store.py @@ -522,12 +522,12 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) else: @@ -538,11 +538,11 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) await cursor.execute(sql, params) @@ -595,6 +595,9 @@ async def search_entries( records: list[StoredMemory] = [] for row in rows: rec = cast("StoredMemory", dict(zip(columns, row, strict=False))) + rec["content_json"] = _json_dict(rec.get("content_json")) + metadata_val = rec.get("metadata_json") + rec["metadata_json"] = _json_dict(metadata_val) if metadata_val is not None else None rec["embedding"] = None records.append(rec) return records diff --git a/sqlspec/adapters/asyncmy/core.py b/sqlspec/adapters/asyncmy/core.py index 906769b7c..536b7d478 100644 --- a/sqlspec/adapters/asyncmy/core.py +++ b/sqlspec/adapters/asyncmy/core.py @@ -90,11 +90,14 @@ async def start(self) -> None: self._driver._check_pending_exception(handler) async def _start(self) -> None: - - cursor = self._driver.connection.cursor(AsyncmySSCursor) - self._cursor = cursor - await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + try: + cursor = self._driver.connection.cursor(AsyncmySSCursor) + self._cursor = cursor + await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) + self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + except BaseException: + await self.close(error=True) + raise async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() diff --git a/sqlspec/adapters/asyncmy/driver.py b/sqlspec/adapters/asyncmy/driver.py index 928c8c95d..eafe6e630 100644 --- a/sqlspec/adapters/asyncmy/driver.py +++ b/sqlspec/adapters/asyncmy/driver.py @@ -446,7 +446,7 @@ def _connection_in_transaction(self) -> bool: Returns: True when the server reports an open transaction. """ - return bool(self.connection.get_transaction_status()) + return bool(getattr(self.connection, "server_status", 0) & 1) def _resolve_column_types(description: Any) -> "dict[str, str] | None": diff --git a/sqlspec/adapters/asyncpg/_typing.py b/sqlspec/adapters/asyncpg/_typing.py index 1368732ea..09c072449 100644 --- a/sqlspec/adapters/asyncpg/_typing.py +++ b/sqlspec/adapters/asyncpg/_typing.py @@ -8,7 +8,8 @@ import asyncpg as asyncpg_module from asyncpg import Connection as AsyncpgRawConnection -from asyncpg import Pool, PostgresError +from asyncpg import Pool +from asyncpg import PostgresError as AsyncpgPostgresError from asyncpg import Record as AsyncpgRecord from asyncpg import connect as asyncpg_connect from asyncpg import create_pool as asyncpg_create_pool @@ -35,13 +36,10 @@ AsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] AsyncpgPool: TypeAlias = Pool[Record] - AsyncpgPostgresError: TypeAlias = PostgresError AsyncpgPreparedStatement: TypeAlias = PreparedStatement[Record] - -if not TYPE_CHECKING: +else: AsyncpgConnection = PoolConnectionProxy AsyncpgPool = Pool - AsyncpgPostgresError = PostgresError AsyncpgPreparedStatement = PreparedStatement diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index 40443d24a..543c15bd3 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -6,6 +6,7 @@ from mypy_extensions import mypyc_attr from typing_extensions import NotRequired +import sqlspec.adapters.asyncpg._typing as _asyncpg_typing from sqlspec.adapters.asyncpg._typing import ( AsyncpgConnection, AsyncpgCursor, @@ -229,7 +230,7 @@ def __init__(self, config: "AsyncpgConfig", user: str | None, password: str | No self._database = database async def __call__(self) -> "AsyncpgConnection": - connector = self._config.get_cloud_sql_connector() + connector = self._config._get_cloud_sql_connector() if connector is None: msg = "Cloud SQL connector is not initialized" raise ImproperConfigurationError(msg) @@ -258,7 +259,7 @@ def __init__(self, config: "AsyncpgConfig", user: str | None, password: str | No self._database = database async def __call__(self) -> "AsyncpgConnection": - connector = self._config.get_alloydb_connector() + connector = self._config._get_alloydb_connector() if connector is None: msg = "AlloyDB connector is not initialized" raise ImproperConfigurationError(msg) @@ -365,11 +366,11 @@ def __init__( self._validate_connector_config() - def get_cloud_sql_connector(self) -> Any | None: + def _get_cloud_sql_connector(self) -> Any | None: """Return the configured Cloud SQL connector instance.""" return self._cloud_sql_connector - def get_alloydb_connector(self) -> Any | None: + def _get_alloydb_connector(self) -> Any | None: """Return the configured AlloyDB connector instance.""" return self._alloydb_connector @@ -426,10 +427,8 @@ def _setup_cloud_sql_connector(self, config: "dict[str, Any]") -> None: Args: config: Pool configuration dictionary to modify in-place. """ - from sqlspec.adapters.asyncpg._typing import AsyncpgCloudSqlConnector as Connector - if self._cloud_sql_connector is None: - self._cloud_sql_connector = Connector() + self._cloud_sql_connector = _asyncpg_typing.AsyncpgCloudSqlConnector() user = config.get("user") password = config.get("password") @@ -446,10 +445,8 @@ def _setup_alloydb_connector(self, config: "dict[str, Any]") -> None: Args: config: Pool configuration dictionary to modify in-place. """ - from sqlspec.adapters.asyncpg._typing import AsyncpgAlloydbAsyncConnector as AsyncConnector - if self._alloydb_connector is None: - self._alloydb_connector = AsyncConnector() + self._alloydb_connector = _asyncpg_typing.AsyncpgAlloydbAsyncConnector() user = config.get("user") password = config.get("password") diff --git a/sqlspec/adapters/bigquery/_typing.py b/sqlspec/adapters/bigquery/_typing.py index 4a3d4b7c3..bef2c4394 100644 --- a/sqlspec/adapters/bigquery/_typing.py +++ b/sqlspec/adapters/bigquery/_typing.py @@ -4,7 +4,7 @@ compilation to avoid ABI boundary issues. """ -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, TypeAlias import google.cloud.bigquery as bigquery_module from google.api_core import exceptions as bigquery_exceptions @@ -22,27 +22,21 @@ from sqlspec.typing import import_optional +BigQueryConnection: TypeAlias = Client +BigQueryParam: TypeAlias = ArrayQueryParameter | ScalarQueryParameter | StructQueryParameter +BigQueryStorageWriteModule: Any = import_optional("google.cloud.bigquery_storage_v1") +BigQueryStorageWriteTypes: Any = import_optional("google.cloud.bigquery_storage_v1.types") + if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias - from google.cloud import bigquery_storage as bigquery_storage_read_module + import google.cloud.bigquery_storage as bigquery_storage_read_module from sqlspec.adapters.bigquery.driver import BigQueryDriver from sqlspec.core import StatementConfig - - BigQueryConnection: TypeAlias = Client - BigQueryParam: TypeAlias = ArrayQueryParameter | ScalarQueryParameter | StructQueryParameter - BigQueryStorageWriteModule: Any - BigQueryStorageWriteTypes: Any - -if not TYPE_CHECKING: - BigQueryConnection = Client - BigQueryParam = ArrayQueryParameter | ScalarQueryParameter | StructQueryParameter +else: bigquery_storage_read_module = import_optional("google.cloud.bigquery_storage") - BigQueryStorageWriteModule = import_optional("google.cloud.bigquery_storage_v1") - BigQueryStorageWriteTypes = import_optional("google.cloud.bigquery_storage_v1.types") __all__ = ( "BIGQUERY_DEFAULT_RETRY", diff --git a/sqlspec/adapters/bigquery/config.py b/sqlspec/adapters/bigquery/config.py index 19fda3587..11072c786 100644 --- a/sqlspec/adapters/bigquery/config.py +++ b/sqlspec/adapters/bigquery/config.py @@ -247,9 +247,6 @@ def __init__( if "default_query_job_config" not in self.connection_config: self._setup_default_job_config() - # Fired directly in create_connection (the client-construction path) like every other adapter, - # rather than bridged through the observability lifecycle dispatcher (which only runs under the - # SQLSpec registry wrapper, not bare config.provide_session()). self._user_connection_hook = user_connection_hook super().__init__( @@ -265,9 +262,9 @@ def __init__( ) self.driver_features = driver_features - self.driver_features["_storage_write_client_provider"] = self.provide_storage_write_client + self.driver_features["_storage_write_client_provider"] = self._provide_storage_write_client - def provide_storage_write_client(self, connection: "BigQueryConnection") -> Any: + def _provide_storage_write_client(self, connection: "BigQueryConnection") -> Any: """Return the shared BigQuery Storage Write client, creating it on first use. Args: diff --git a/sqlspec/adapters/bigquery/core.py b/sqlspec/adapters/bigquery/core.py index 1ea8205f1..b45bb6156 100644 --- a/sqlspec/adapters/bigquery/core.py +++ b/sqlspec/adapters/bigquery/core.py @@ -12,7 +12,14 @@ import sqlglot from sqlglot import exp +from sqlspec.adapters.bigquery._typing import BIGQUERY_DEFAULT_RETRY as DEFAULT_RETRY +from sqlspec.adapters.bigquery._typing import BigQueryConnection, BigQueryParam, GoogleCloudError +from sqlspec.adapters.bigquery._typing import BigQueryLoadJobConfig as LoadJobConfig +from sqlspec.adapters.bigquery._typing import BigQueryQueryJob as QueryJob +from sqlspec.adapters.bigquery._typing import BigQueryQueryJobConfig as QueryJobConfig +from sqlspec.adapters.bigquery._typing import BigQueryRetry as Retry from sqlspec.adapters.bigquery._typing import bigquery_exceptions as api_exceptions +from sqlspec.adapters.bigquery._typing import bigquery_module as bigquery from sqlspec.core import ( DriverParameterProfile, ParameterProfile, @@ -45,11 +52,6 @@ from collections.abc import Callable, Iterable, Iterator, Mapping from typing import Literal - from sqlspec.adapters.bigquery._typing import BigQueryConnection, BigQueryParam - from sqlspec.adapters.bigquery._typing import BigQueryLoadJobConfig as LoadJobConfig - from sqlspec.adapters.bigquery._typing import BigQueryQueryJob as QueryJob - from sqlspec.adapters.bigquery._typing import BigQueryQueryJobConfig as QueryJobConfig - from sqlspec.adapters.bigquery._typing import BigQueryRetry as Retry from sqlspec.driver import SyncExceptionHandler from sqlspec.storage import StorageFormat, StorageTelemetry from sqlspec.typing import StatementParameters @@ -283,8 +285,6 @@ def create_parameters(parameters: Any, json_serializer: "Callable[[Any], str] | def build_retry(deadline: float) -> "Retry": """Build retry policy for job restarts based on error reason codes.""" - from sqlspec.adapters.bigquery._typing import BigQueryRetry as Retry - return Retry(predicate=_should_retry_bigquery_job, deadline=deadline) @@ -365,8 +365,6 @@ def run_query_job( Returns: QueryJob object representing the executed job. """ - from sqlspec.adapters.bigquery._typing import BigQueryQueryJobConfig as QueryJobConfig - final_job_config = QueryJobConfig() if default_job_config: copy_job_config(default_job_config, final_job_config) @@ -392,8 +390,6 @@ def run_query_job( def build_load_job_config(file_format: "BigQueryLoadFormat", overwrite: bool) -> "LoadJobConfig": - from sqlspec.adapters.bigquery._typing import BigQueryLoadJobConfig as LoadJobConfig - job_config = LoadJobConfig() job_config.source_format = _map_bigquery_source_format(file_format) job_config.write_disposition = "WRITE_TRUNCATE" if overwrite else "WRITE_APPEND" @@ -525,8 +521,6 @@ def __init__( self._pages: Iterator[Iterable[_BigQueryRow]] | None = None def start(self) -> None: - from sqlspec.adapters.bigquery._typing import BIGQUERY_DEFAULT_RETRY as DEFAULT_RETRY - handler = self._driver.handle_database_exceptions() with handler: page_size = None if _uses_local_bigquery_endpoint(self._driver.connection) else self._chunk_size @@ -999,12 +993,7 @@ def _is_query_parameter(value: Any) -> bool: def _load_bigquery_module() -> Any: - global _BIGQUERY_MODULE - if _BIGQUERY_MODULE is None: - from sqlspec.adapters.bigquery._typing import bigquery_module as bigquery - - _BIGQUERY_MODULE = bigquery - return _BIGQUERY_MODULE + return bigquery if _BIGQUERY_MODULE is None else _BIGQUERY_MODULE def _query_parameter_type(value: Any, declared_type: "type[Any] | None" = None) -> "tuple[str | None, str | None]": @@ -1065,8 +1054,6 @@ def _inline_bigquery_literals( def _should_retry_bigquery_job(exception: Exception) -> bool: """Return True when a BigQuery job exception is safe to retry.""" - from sqlspec.adapters.bigquery._typing import GoogleCloudError - if not isinstance(exception, GoogleCloudError): return False @@ -1113,8 +1100,6 @@ def _run_query_and_wait( max_results: int | None = None, ) -> Any: """Execute a BigQuery query via query_and_wait and return the row iterator.""" - from sqlspec.adapters.bigquery._typing import BigQueryQueryJobConfig as QueryJobConfig - final_job_config = QueryJobConfig() if default_job_config: copy_job_config(default_job_config, final_job_config) diff --git a/sqlspec/adapters/bigquery/driver.py b/sqlspec/adapters/bigquery/driver.py index 22a72dadc..b2db2eb76 100644 --- a/sqlspec/adapters/bigquery/driver.py +++ b/sqlspec/adapters/bigquery/driver.py @@ -59,7 +59,7 @@ from sqlspec.driver import BaseSyncExceptionHandler, ExecutionResult, SyncDriverAdapterBase, SyncRowStream from sqlspec.exceptions import ImproperConfigurationError, StorageOperationFailedError from sqlspec.utils.logging import get_logger -from sqlspec.utils.module_loader import ensure_pyarrow +from sqlspec.utils.module_loader import ensure_pyarrow, import_optional from sqlspec.utils.serializers import to_json from sqlspec.utils.text import split_qualified_identifier @@ -81,6 +81,8 @@ logger = get_logger(__name__) _DATASET_TABLE_PARTS = 2 _PROJECT_DATASET_TABLE_PARTS = 3 +_pa: Any = import_optional("pyarrow") +_pq: Any = import_optional("pyarrow.parquet") class BigQueryExceptionHandler(BaseSyncExceptionHandler): @@ -297,12 +299,15 @@ def dispatch_execute_many(self, cursor: Any, statement: "SQL") -> ExecutionResul def dispatch_execute_script(self, cursor: Any, statement: "SQL") -> ExecutionResult: """Execute a procedural script as one native BigQuery job with bound parameters.""" sql, parameters = self._compiled_sql(statement, self.statement_config) + stmt_count = max( + len(self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True)), 1 + ) cursor.job = self._run_query_job(cursor, sql, parameters) cursor.job.result(job_retry=self._job_retry, timeout=self._job_result_timeout) return self.create_execution_result( cursor, - statement_count=1, - successful_statements=1, + statement_count=stmt_count, + successful_statements=stmt_count, rowcount_override=normalize_script_rowcount(0, cursor.job), is_script_result=True, ) @@ -470,9 +475,7 @@ def select_to_arrow( with exc_handler: query_job = self._run_query_job(self.connection, sql, driver_params) - query_job.result( - job_retry=self._job_retry, timeout=self._job_result_timeout, **self._job_result_kwargs() - ) # Wait for completion + query_job.result(job_retry=self._job_retry, timeout=self._job_result_timeout, **self._job_result_kwargs()) arrow_table = query_job.to_arrow() @@ -607,10 +610,8 @@ def load_from_arrow( self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload) - import pyarrow.parquet as pq - buffer = io.BytesIO() - pq.write_table(arrow_table, buffer) + _pq.write_table(arrow_table, buffer) buffer.seek(0) job_config = build_load_job_config("parquet", overwrite) job = self.connection.load_table_from_file( @@ -829,14 +830,15 @@ def _resolve_storage_write_table_path(table: str, default_project: str) -> tuple def _bigquery_arrow_reader_from_iterable(batches: "Iterable[ArrowRecordBatch]") -> "ArrowRecordBatchReader | None": ensure_pyarrow() - import pyarrow as pa - iterator = iter(batches) try: first_batch = next(iterator) except StopIteration: return None - return pa.RecordBatchReader.from_batches(first_batch.schema, chain((first_batch,), iterator)) + return cast( + "ArrowRecordBatchReader", + _pa.RecordBatchReader.from_batches(first_batch.schema, chain((first_batch,), iterator)), + ) def _records_to_json_rows( diff --git a/sqlspec/adapters/bigquery/litestar/store.py b/sqlspec/adapters/bigquery/litestar/store.py index 1fb07ec54..c63e0858f 100644 --- a/sqlspec/adapters/bigquery/litestar/store.py +++ b/sqlspec/adapters/bigquery/litestar/store.py @@ -174,6 +174,20 @@ def _drop_table_sql(self) -> "list[str]": """ return [f"DROP TABLE IF EXISTS {self._table_name}"] + def _partition_filter(self, qualifier: str = "") -> str: + """Return a constant partition predicate when require_partition_filter is enabled. + + Args: + qualifier: Optional table alias qualifier for the expires_at column. + + Returns: + SQL predicate fragment or empty string when not required. + """ + if not self._require_partition_filter: + return "" + prefix = f"{qualifier}." if qualifier else "" + return f" AND ({prefix}expires_at IS NULL OR {prefix}expires_at >= TIMESTAMP('1970-01-01 00:00:00+00'))" + def _datetime_to_timestamp(self, dt: "datetime | None") -> "datetime | None": """Convert datetime to BigQuery TIMESTAMP. @@ -235,7 +249,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | update_sql = f""" UPDATE {self._table_name} SET expires_at = @expires_at - WHERE session_id = @session_id + WHERE session_id = @session_id{self._partition_filter()} """ driver.execute(update_sql, expires_at=new_expires_at_ts, session_id=key) @@ -250,7 +264,7 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No sql = f""" MERGE {self._table_name} AS target USING (SELECT @session_id AS session_id, @data AS data, @expires_at AS expires_at) AS source - ON target.session_id = source.session_id + ON target.session_id = source.session_id{self._partition_filter("target")} WHEN MATCHED THEN UPDATE SET data = source.data, expires_at = source.expires_at WHEN NOT MATCHED THEN @@ -263,14 +277,14 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No def _delete(self, key: str) -> None: """Synchronous implementation of delete.""" - sql = f"DELETE FROM {self._table_name} WHERE session_id = @session_id" + sql = f"DELETE FROM {self._table_name} WHERE session_id = @session_id{self._partition_filter()}" with self._config.provide_session() as driver: driver.execute(sql, session_id=key) def _delete_all(self) -> None: """Synchronous implementation of delete_all.""" - sql = f"DELETE FROM {self._table_name} WHERE TRUE" + sql = f"DELETE FROM {self._table_name} WHERE TRUE{self._partition_filter()}" with self._config.provide_session() as driver: driver.execute(sql) @@ -293,7 +307,7 @@ def _expires_in(self, key: str) -> "int | None": """Synchronous implementation of expires_in.""" sql = f""" SELECT expires_at FROM {self._table_name} - WHERE session_id = @session_id + WHERE session_id = @session_id{self._partition_filter()} """ with self._config.provide_session() as driver: diff --git a/sqlspec/adapters/cockroach_asyncpg/_typing.py b/sqlspec/adapters/cockroach_asyncpg/_typing.py index d3362b6c6..c71e745c8 100644 --- a/sqlspec/adapters/cockroach_asyncpg/_typing.py +++ b/sqlspec/adapters/cockroach_asyncpg/_typing.py @@ -3,7 +3,8 @@ from typing import TYPE_CHECKING, Any import asyncpg as cockroach_asyncpg_module -from asyncpg import Pool, PostgresError +from asyncpg import Pool +from asyncpg import PostgresError as CockroachAsyncpgPostgresError from asyncpg import Record as CockroachAsyncpgRecord from asyncpg import connect as cockroach_asyncpg_connect from asyncpg import create_pool as cockroach_asyncpg_create_pool @@ -20,12 +21,9 @@ from sqlspec.core import StatementConfig CockroachAsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] - CockroachAsyncpgPostgresError: TypeAlias = PostgresError CockroachAsyncpgPool: TypeAlias = Pool[Record] - -if not TYPE_CHECKING: +else: CockroachAsyncpgConnection = PoolConnectionProxy - CockroachAsyncpgPostgresError = PostgresError CockroachAsyncpgPool = Pool __all__ = ( diff --git a/sqlspec/adapters/cockroach_psycopg/_typing.py b/sqlspec/adapters/cockroach_psycopg/_typing.py index eead89a5e..6a6e52624 100644 --- a/sqlspec/adapters/cockroach_psycopg/_typing.py +++ b/sqlspec/adapters/cockroach_psycopg/_typing.py @@ -9,7 +9,6 @@ import psycopg as cockroach_psycopg_module from psycopg import AsyncCursor, Cursor from psycopg import crdb as cockroach_psycopg_crdb -from psycopg import crdb as psycopg_crdb from psycopg import errors as cockroach_psycopg_errors from psycopg import sql as cockroach_psycopg_sql from psycopg.rows import DictRow as PsycopgDictRow @@ -32,10 +31,9 @@ CockroachAsyncConnection: TypeAlias = AsyncCrdbConnection[PsycopgDictRow] CockroachSyncCursor: TypeAlias = Cursor[PsycopgDictRow] CockroachAsyncCursor: TypeAlias = AsyncCursor[PsycopgDictRow] - -if not TYPE_CHECKING: - CockroachSyncConnection = psycopg_crdb.CrdbConnection - CockroachAsyncConnection = psycopg_crdb.AsyncCrdbConnection +else: + CockroachSyncConnection = cockroach_psycopg_crdb.CrdbConnection + CockroachAsyncConnection = cockroach_psycopg_crdb.AsyncCrdbConnection CockroachSyncCursor = Cursor CockroachAsyncCursor = AsyncCursor diff --git a/sqlspec/adapters/cockroach_psycopg/adk/store.py b/sqlspec/adapters/cockroach_psycopg/adk/store.py index 8d81ec748..0d2ff1781 100644 --- a/sqlspec/adapters/cockroach_psycopg/adk/store.py +++ b/sqlspec/adapters/cockroach_psycopg/adk/store.py @@ -140,6 +140,7 @@ " session_id VARCHAR(128) NOT NULL,\n" " app_name VARCHAR(128) NOT NULL,\n" " user_id VARCHAR(128) NOT NULL,\n" + " scope VARCHAR(16) NOT NULL DEFAULT 'user',\n" " event_id VARCHAR(128) NOT NULL UNIQUE,\n" " author VARCHAR(256){1},\n" " timestamp TIMESTAMPTZ NOT NULL,\n" @@ -261,7 +262,8 @@ async def create_session( async def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None ) -> "StoredSession | None": - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + should_touch = renew_for is not None and self._calculate_expires_at(renew_for) is not None + if should_touch: sql = f""" UPDATE {self._session_table} SET update_time = CURRENT_TIMESTAMP @@ -279,6 +281,8 @@ async def get_session( async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(sql.encode(), (app_name, user_id, session_id)) row = await cur.fetchone() + if should_touch: + await conn.commit() if row is None: return None @@ -748,7 +752,8 @@ def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None ) -> "StoredSession | None": """Get session by ID.""" - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + should_touch = renew_for is not None and self._calculate_expires_at(renew_for) is not None + if should_touch: sql = f""" UPDATE {self._session_table} SET update_time = CURRENT_TIMESTAMP @@ -766,6 +771,8 @@ def get_session( with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute(sql.encode(), (app_name, user_id, session_id)) row = cur.fetchone() + if should_touch: + conn.commit() if row is None: return None @@ -841,7 +848,6 @@ def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: def append_event(self, event_record: StoredEvent) -> None: """Append an event to a session.""" - """Synchronous implementation of append_event.""" self._insert_event(event_record) def append_event_and_update_state( @@ -1435,6 +1441,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object cur.execute(query, _build_insert_params(entry)) if cur.rowcount and cur.rowcount > 0: inserted_count += cur.rowcount + conn.commit() return inserted_count def search_entries( diff --git a/sqlspec/adapters/cockroach_psycopg/driver.py b/sqlspec/adapters/cockroach_psycopg/driver.py index 7871e1527..d5421d567 100644 --- a/sqlspec/adapters/cockroach_psycopg/driver.py +++ b/sqlspec/adapters/cockroach_psycopg/driver.py @@ -39,7 +39,8 @@ from collections.abc import Awaitable, Callable from sqlspec.adapters.cockroach_psycopg._typing import CockroachAsyncCursor, CockroachSyncCursor - from sqlspec.driver import ExecutionResult + from sqlspec.core import SQLResult + from sqlspec.driver import CachedQuery, ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry __all__ = ( @@ -105,9 +106,15 @@ def __init__( self._retry_config = CockroachPsycopgRetryConfig.from_features(self.driver_features) self._enable_retry = bool(self.driver_features.get("enable_auto_retry", True)) self._follower_staleness = cast("str | None", self.driver_features.get("default_staleness")) - # Data dictionary is lazily initialized in property; use parent slot self._data_dictionary = None + def _execute_cache_hit( + self, sql: str, params: "tuple[Any, ...] | list[Any] | dict[str, Any]", cached: "CachedQuery" + ) -> "SQLResult": + if cached.operation_profile.returns_rows: + self._begin_follower_read_transaction() + return super()._execute_cache_hit(sql, params, cached) + def select_to_storage( self, statement: "SQL | str", @@ -275,7 +282,6 @@ def handle_database_exceptions(self) -> "CockroachPsycopgSyncExceptionHandler": @property def data_dictionary(self) -> "CockroachPsycopgSyncDataDictionary": # type: ignore[override] if self._data_dictionary is None: - # Intentionally assign CockroachDB-specific data dictionary to parent slot self._data_dictionary = CockroachPsycopgSyncDataDictionary() # type: ignore[assignment] return cast("CockroachPsycopgSyncDataDictionary", self._data_dictionary) @@ -336,9 +342,15 @@ def __init__( self._retry_config = CockroachPsycopgRetryConfig.from_features(self.driver_features) self._enable_retry = bool(self.driver_features.get("enable_auto_retry", True)) self._follower_staleness = cast("str | None", self.driver_features.get("default_staleness")) - # Data dictionary is lazily initialized in property; use parent slot self._data_dictionary = None + async def _execute_cache_hit( + self, sql: str, params: "tuple[Any, ...] | list[Any] | dict[str, Any]", cached: "CachedQuery" + ) -> "SQLResult": + if cached.operation_profile.returns_rows: + await self._begin_follower_read_transaction() + return await super()._execute_cache_hit(sql, params, cached) + async def select_to_storage( self, statement: "SQL | str", @@ -506,7 +518,6 @@ def handle_database_exceptions(self) -> "CockroachPsycopgAsyncExceptionHandler": @property def data_dictionary(self) -> "CockroachPsycopgAsyncDataDictionary": # type: ignore[override] if self._data_dictionary is None: - # Intentionally assign CockroachDB-specific data dictionary to parent slot self._data_dictionary = CockroachPsycopgAsyncDataDictionary() # type: ignore[assignment] return cast("CockroachPsycopgAsyncDataDictionary", self._data_dictionary) diff --git a/sqlspec/adapters/db2/_typing.py b/sqlspec/adapters/db2/_typing.py index b7edaf0c2..9c7f3d458 100644 --- a/sqlspec/adapters/db2/_typing.py +++ b/sqlspec/adapters/db2/_typing.py @@ -26,10 +26,10 @@ class _Db2UnavailableError(Exception): from sqlspec.core import StatementConfig if not TYPE_CHECKING: - Db2SyncConnection = import_optional_attr("ibm_db_dbi", "Connection") or Any - Db2RawCursor = import_optional_attr("ibm_db_dbi", "Cursor") or Any - Db2AsyncConnection = import_optional_attr("ibm_db_dbi", "AsyncConnection") or Any - Db2AsyncRawCursor = import_optional_attr("ibm_db_dbi", "AsyncCursor") or Any + Db2SyncConnection = Any + Db2RawCursor = Any + Db2AsyncConnection = Any + Db2AsyncRawCursor = Any Db2Error = import_optional_attr("ibm_db_dbi", "Error") or _Db2UnavailableError ibm_db = import_optional("ibm_db") @@ -183,7 +183,7 @@ def __exit__( return None try: if self._driver is not None: - self._driver.release_open_work(autocommit_baseline=self._autocommit_baseline) + self._driver._release_open_work(autocommit_baseline=self._autocommit_baseline) finally: self._release_connection(self._connection, exc_type=exc_type, exc_val=exc_val, exc_tb=exc_tb) self._connection = None @@ -303,7 +303,7 @@ async def __aexit__( return None try: if self._driver is not None: - await self._driver.release_open_work(autocommit_baseline=self._autocommit_baseline) + await self._driver._release_open_work(autocommit_baseline=self._autocommit_baseline) finally: await self._release_connection(self._connection, exc_type=exc_type, exc_val=exc_val, exc_tb=exc_tb) self._connection = None diff --git a/sqlspec/adapters/db2/driver.py b/sqlspec/adapters/db2/driver.py index 3e678a968..c60f9b0ca 100644 --- a/sqlspec/adapters/db2/driver.py +++ b/sqlspec/adapters/db2/driver.py @@ -250,10 +250,11 @@ def rollback(self) -> None: except Db2Error as exc: msg = f"Failed to rollback Db2 transaction: {exc}" raise SQLSpecError(msg) from exc - self._transaction_active = False - self._restore_connection_autocommit() + finally: + self._transaction_active = False + self._restore_connection_autocommit() - def release_open_work(self, *, autocommit_baseline: bool) -> None: + def _release_open_work(self, *, autocommit_baseline: bool) -> None: """Roll back work left open before the connection is returned to its pool. A transaction started by this driver is always rolled back; on a connection whose @@ -554,10 +555,11 @@ async def rollback(self) -> None: except Db2Error as exc: msg = f"Failed to rollback Db2 transaction: {exc}" raise SQLSpecError(msg) from exc - self._transaction_active = False - await self._restore_connection_autocommit() + finally: + self._transaction_active = False + await self._restore_connection_autocommit() - async def release_open_work(self, *, autocommit_baseline: bool) -> None: + async def _release_open_work(self, *, autocommit_baseline: bool) -> None: """Roll back work left open before the connection is returned to its pool. A transaction started by this driver is always rolled back; on a connection whose diff --git a/sqlspec/adapters/db2/pool.py b/sqlspec/adapters/db2/pool.py index 37a4e9178..1e4ff74d4 100644 --- a/sqlspec/adapters/db2/pool.py +++ b/sqlspec/adapters/db2/pool.py @@ -252,13 +252,9 @@ def release(self, connection: Any) -> None: _ = connection def size(self) -> int: - """Return the count of active connections allocated to the current thread.""" - try: - _ = self._thread_local.connection - except AttributeError: - return 0 - else: - return 1 + """Return the total number of active connections registered across all threads.""" + with self._registry_lock: + return len(self._connection_registry) def checked_out(self) -> int: """Return the number of checked out connections from the perspective of this thread.""" diff --git a/sqlspec/adapters/duckdb/_typing.py b/sqlspec/adapters/duckdb/_typing.py index f1453cb89..ab6d8232b 100644 --- a/sqlspec/adapters/duckdb/_typing.py +++ b/sqlspec/adapters/duckdb/_typing.py @@ -29,8 +29,7 @@ from sqlspec.core import StatementConfig DuckDBConnection: TypeAlias = _DuckDBConnection - -if not TYPE_CHECKING: +else: DuckDBConnection = _DuckDBConnection __all__ = ( diff --git a/sqlspec/adapters/duckdb/adk/store.py b/sqlspec/adapters/duckdb/adk/store.py index ef4ff5a2e..44b127880 100644 --- a/sqlspec/adapters/duckdb/adk/store.py +++ b/sqlspec/adapters/duckdb/adk/store.py @@ -946,6 +946,9 @@ def __init__(self, config: "DuckDBConfig") -> None: def create_tables(self) -> None: """Create the memory table and indexes if they don't exist.""" + if not self._enabled: + return + if not self.create_schema_enabled: self.reconcile_schema() return @@ -1131,6 +1134,7 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec for entry in entries: params: tuple[Any, ...] scope_value = entry.get("scope", "user") + metadata_json = entry.get("metadata_json") if self._owner_id_column_name: params = ( entry["id"], @@ -1139,12 +1143,12 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], scope_value, entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(metadata_json) if metadata_json is not None else None, entry["inserted_at"], ) else: @@ -1155,18 +1159,17 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], scope_value, entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(metadata_json) if metadata_json is not None else None, entry["inserted_at"], ) result = conn.execute(sql, params) inserted_count += len(result.fetchall()) conn.commit() - # Refresh FTS index after inserts, not on search if self._use_fts and inserted_count > 0: self._refresh_fts_index(conn) @@ -1185,20 +1188,21 @@ def _search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) - if not query: + if not query or not query.strip(): return [] - limit_value = limit or self._max_results + limit_value = limit if limit is not None else self._max_results use_fts = self._use_fts + rows: list[Any] | None = None + columns: list[str] = [] with self._config.provide_connection() as conn: if use_fts and not self._ensure_fts_extension(conn): use_fts = False if use_fts: - # Use match_bm25() -- the correct DuckDB FTS syntax where_scope, scope_params = _build_duckdb_scope_where(app_name, user_id, scope_filter, prefix="m") - sql = f""" + fts_sql = f""" SELECT m.* FROM {self._memory_table} m JOIN ( @@ -1209,8 +1213,14 @@ def _search_entries( ORDER BY fts.score DESC LIMIT ? """ - params = (query, *scope_params, limit_value) - else: + fts_params = (query, *scope_params, limit_value) + try: + rows = conn.execute(fts_sql, fts_params).fetchall() + columns = [col[0] for col in conn.description or []] + except Exception as exc: + logger.warning("FTS search failed; falling back to simple search: %s", exc) + + if rows is None: where_scope, scope_params = _build_duckdb_scope_where(app_name, user_id, scope_filter) sql = f""" SELECT * FROM {self._memory_table} @@ -1219,9 +1229,14 @@ def _search_entries( LIMIT ? """ params = (*scope_params, f"%{query}%", limit_value) + try: + rows = conn.execute(sql, params).fetchall() + columns = [col[0] for col in conn.description or []] + except Exception as exc: + if DUCKDB_TABLE_NOT_FOUND_ERROR in str(exc): + return [] + raise - rows = conn.execute(sql, params).fetchall() - columns = [col[0] for col in conn.description or []] records: list[StoredMemory] = [] for row in rows: record = cast("StoredMemory", dict(zip(columns, row, strict=False))) @@ -1256,8 +1271,8 @@ def _delete_entries_older_than(self, days: int, app_name: "str | None" = None, s msg = "Memory store is disabled" raise RuntimeError(msg) - clauses = [f"inserted_at < (CURRENT_TIMESTAMP - INTERVAL '{days} days')"] - params: list[Any] = [] + clauses = ["inserted_at < (CURRENT_TIMESTAMP - (? * INTERVAL '1 day'))"] + params: list[Any] = [days] if app_name is not None: clauses.append("app_name = ?") params.append(app_name) diff --git a/sqlspec/adapters/duckdb/driver.py b/sqlspec/adapters/duckdb/driver.py index 77d2cd9b1..9eeeae7f6 100644 --- a/sqlspec/adapters/duckdb/driver.py +++ b/sqlspec/adapters/duckdb/driver.py @@ -35,6 +35,7 @@ from sqlspec.core.result import DMLResult from sqlspec.driver import BaseSyncExceptionHandler, SyncDriverAdapterBase, SyncRowStream from sqlspec.exceptions import SQLSpecError +from sqlspec.typing import import_optional from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import ensure_pyarrow from sqlspec.utils.text import quote_identifier @@ -43,13 +44,16 @@ if TYPE_CHECKING: from collections.abc import Sequence + import pyarrow as pa + from sqlspec.adapters.duckdb._typing import DuckDBConnection from sqlspec.builder import QueryBuilder from sqlspec.core import ArrowResult, SQLResult, Statement, StatementFilter from sqlspec.driver import ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import ArrowReturnFormat, StatementParameters - +else: + pa = import_optional("pyarrow") __all__ = ("DuckDBCursor", "DuckDBDriver", "DuckDBExceptionHandler", "DuckDBSessionContext") @@ -147,7 +151,8 @@ def dispatch_execute(self, cursor: "DuckDBConnection", statement: SQL) -> "Execu if is_select_like: arrow_table = cursor.to_arrow_table() data = arrow_table.to_pylist() - _restore_uuid_columns(data, cursor.description) + if self.driver_features.get("enable_uuid_conversion", True): + _restore_uuid_columns(data, cursor.description) column_names = list(arrow_table.column_names) return self.create_execution_result( @@ -465,7 +470,6 @@ def load_from_arrow( self._require_capability("arrow_import_enabled") ensure_pyarrow() - import pyarrow as pa source_data = source.get_data() if hasattr(source, "get_data") else source arrow_table = None @@ -628,19 +632,22 @@ def _execute_bulk_insert_many(self, expression: exp.Insert, prepared_parameters: target_table = table_expr.sql(dialect="duckdb") column_sql = ", ".join(column.sql(dialect="duckdb") for column in expression.this.expressions) temp_view = f"_sqlspec_batch_{uuid4().hex}" - self.connection.register(temp_view, arrow_table) - try: - self.connection.execute(f"INSERT INTO {target_table} ({column_sql}) SELECT * FROM {temp_view}") - finally: - with contextlib.suppress(Exception): - self.connection.unregister(temp_view) + exc_handler = self.handle_database_exceptions() + with exc_handler: + self.connection.register(temp_view, arrow_table) + try: + self.connection.execute(f"INSERT INTO {target_table} ({column_sql}) SELECT * FROM {temp_view}") + finally: + with contextlib.suppress(Exception): + self.connection.unregister(temp_view) + self._check_pending_exception(exc_handler) return DMLResult("INSERT", len(rows)) @staticmethod def _build_arrow_table(rows: "list[Any]", column_names: "list[str]") -> Any | None: """Build a pyarrow table from batch rows when they share a stable shape.""" - if not rows: + if not rows or pa is None: return None first_row = rows[0] @@ -648,9 +655,10 @@ def _build_arrow_table(rows: "list[Any]", column_names: "list[str]") -> Any | No keys = column_names or list(first_row.keys()) if any(not isinstance(row, dict) for row in rows): return None - import pyarrow as pa - - return pa.table({key: [row.get(key) for row in rows] for key in keys}) + try: + return pa.table({key: [row.get(key) for row in rows] for key in keys}) + except Exception: + return None if isinstance(first_row, (list, tuple)): values = list(first_row) @@ -658,9 +666,10 @@ def _build_arrow_table(rows: "list[Any]", column_names: "list[str]") -> Any | No column_names = [f"col_{index}" for index in range(len(values))] if any(not isinstance(row, (list, tuple)) or len(row) != len(column_names) for row in rows): return None - import pyarrow as pa - - return pa.Table.from_arrays([pa.array(col) for col in zip(*rows, strict=True)], names=column_names) + try: + return pa.Table.from_arrays([pa.array(col) for col in zip(*rows, strict=True)], names=column_names) + except Exception: + return None return None @@ -683,7 +692,7 @@ def _open_stream_reader(self, sql: str, parameters: Any, chunk_size: int) -> "tu reader: Any | None = None with handler: result = self.connection.execute(sql, normalize_execute_parameters(parameters)) - description = result.description + description = result.description if self.driver_features.get("enable_uuid_conversion", True) else None reader = result.to_arrow_reader(chunk_size) self._check_pending_exception(handler) if reader is None: diff --git a/sqlspec/adapters/duckdb/litestar/store.py b/sqlspec/adapters/duckdb/litestar/store.py index 01abe10a8..025ea3163 100644 --- a/sqlspec/adapters/duckdb/litestar/store.py +++ b/sqlspec/adapters/duckdb/litestar/store.py @@ -291,12 +291,12 @@ def _expires_in(self, key: str) -> "int | None": def _delete_expired(self) -> int: """Synchronous implementation of delete_expired.""" - sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP" + sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP RETURNING 1" with self._config.provide_connection() as conn: cursor = conn.execute(sql) - count = cursor.fetchone() - row_count = count[0] if count else 0 + row_count = len(cursor.fetchall()) + conn.commit() if row_count > 0: self._log_delete_expired(row_count) return row_count diff --git a/sqlspec/adapters/mssql_python/adk/store.py b/sqlspec/adapters/mssql_python/adk/store.py index 1626a210a..7aacf57c1 100644 --- a/sqlspec/adapters/mssql_python/adk/store.py +++ b/sqlspec/adapters/mssql_python/adk/store.py @@ -7,6 +7,7 @@ from sqlspec.adapters.mssql_python._typing import MssqlPythonCursor, MssqlPythonError from sqlspec.adapters.mssql_python.core import extract_error_number +from sqlspec.adapters.mssql_python.data_dictionary import MssqlVersionInfo from sqlspec.config import ADKConfig from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options from sqlspec.extensions.adk.memory.store import BaseSyncADKMemoryStore @@ -17,6 +18,7 @@ from datetime import timedelta from sqlspec.adapters.mssql_python.config import MssqlPythonConfig + from sqlspec.adapters.mssql_python.driver import MssqlPythonDriver from sqlspec.extensions.adk import SessionOrderBy from sqlspec.extensions.adk.memory._types import StoredMemory @@ -41,13 +43,14 @@ class MssqlPythonADKStore(BaseSyncADKStore["MssqlPythonConfig"]): """Synchronous mssql-python ADK session/event store.""" connector_name: ClassVar[str] = "mssql_python" - __slots__ = ("_json_column_type",) + __slots__ = ("_json_column_type", "_native_json") def __init__(self, config: "MssqlPythonConfig") -> None: super().__init__(config) adk_config = _adk_config(config) native_json = adk_config.get("native_json") - self._json_column_type = JSON_NATIVE_COLUMN_TYPE if native_json is True else JSON_FALLBACK_COLUMN_TYPE + self._native_json: bool | None = native_json if isinstance(native_json, bool) else None + self._json_column_type: str | None = None def create_tables(self) -> None: """Create ADK tables (idempotent T-SQL) and DD-gated indexes.""" @@ -56,6 +59,11 @@ def create_tables(self) -> None: return with self._config.provide_session() as driver: + if self._json_column_type is None: + configured = _configured_json_column_type(self._native_json) + self._json_column_type = ( + configured if configured is not None else _json_column_type_from_sync_driver(driver) + ) driver.execute_script(self._sessions_table_ddl()) driver.execute_script(self._events_table_ddl()) driver.execute_script(self._app_states_table_ddl()) @@ -381,6 +389,17 @@ def _events_query( return _events_query(self._events_table, app_name, user_id, session_id, after_timestamp, limit) def _json_column_type_sync(self) -> str: + if self._json_column_type is not None: + return self._json_column_type + configured = _configured_json_column_type(self._native_json) + if configured is not None: + self._json_column_type = configured + return configured + try: + with self._config.provide_session() as driver: + self._json_column_type = _json_column_type_from_sync_driver(driver) + except Exception: + return JSON_FALLBACK_COLUMN_TYPE return self._json_column_type def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": @@ -441,7 +460,6 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", ?" if self._owner_id_column_name else "" - # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -467,7 +485,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry.get("metadata_json")), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, ) if self._owner_id_column_name: params = (*params, owner_id) @@ -589,6 +607,21 @@ def _adk_config(config: Any) -> MssqlPythonADKConfig: return cast("MssqlPythonADKConfig", adk_config) +def _configured_json_column_type(native_json: "bool | None") -> "str | None": + if native_json is None: + return None + if native_json is True: + return JSON_NATIVE_COLUMN_TYPE + return JSON_FALLBACK_COLUMN_TYPE + + +def _json_column_type_from_sync_driver(driver: "MssqlPythonDriver") -> str: + version_info = driver.data_dictionary.get_version(driver) + if isinstance(version_info, MssqlVersionInfo) and version_info.supports_native_json(): + return JSON_NATIVE_COLUMN_TYPE + return JSON_FALLBACK_COLUMN_TYPE + + def _sessions_table_ddl(table: str, json_column_type: str, owner_id_column_ddl: "str | None") -> str: owner_line = f",\n {owner_id_column_ddl}" if owner_id_column_ddl else "" return f""" diff --git a/sqlspec/adapters/mssql_python/config.py b/sqlspec/adapters/mssql_python/config.py index 0fda4dc80..d719410a6 100644 --- a/sqlspec/adapters/mssql_python/config.py +++ b/sqlspec/adapters/mssql_python/config.py @@ -227,6 +227,7 @@ def _create_pool(self) -> "MssqlPythonConnectionPool": def _close_pool(self) -> None: if self.connection_instance is not None: self.connection_instance.close() + self.connection_instance = None def _apply_json_serializer_override(statement_config: Any, features_dict: dict[str, Any]) -> Any: diff --git a/sqlspec/adapters/mssql_python/core.py b/sqlspec/adapters/mssql_python/core.py index dda0049d8..4811f9e8c 100644 --- a/sqlspec/adapters/mssql_python/core.py +++ b/sqlspec/adapters/mssql_python/core.py @@ -105,7 +105,7 @@ def extract_error_number(exc: BaseException | None) -> int | None: return None for attr in ("number", "error_code", "errno"): val = getattr(exc, attr, None) - if isinstance(val, int) and not isinstance(val, bool): + if isinstance(val, int) and not isinstance(val, bool) and val != 0: return val ddbc_err = getattr(exc, "ddbc_error", None) if isinstance(ddbc_err, str) and ddbc_err.startswith("("): @@ -121,7 +121,7 @@ def extract_error_number(exc: BaseException | None) -> int | None: if exc.args: first_arg = exc.args[0] - if isinstance(first_arg, int) and not isinstance(first_arg, bool): + if isinstance(first_arg, int) and not isinstance(first_arg, bool) and first_arg != 0: return first_arg if isinstance(first_arg, str): msg = first_arg diff --git a/sqlspec/adapters/mssql_python/driver.py b/sqlspec/adapters/mssql_python/driver.py index c35982958..fb152a780 100644 --- a/sqlspec/adapters/mssql_python/driver.py +++ b/sqlspec/adapters/mssql_python/driver.py @@ -240,8 +240,9 @@ def rollback(self) -> None: except MssqlPythonError as exc: msg = f"Failed to rollback transaction: {exc}" raise SQLSpecError(msg) from exc - self._transaction_active = False - self._restore_connection_autocommit() + finally: + self._transaction_active = False + self._restore_connection_autocommit() def with_cursor(self, connection: "MssqlPythonConnection") -> "MssqlPythonCursor": return MssqlPythonCursor(connection) diff --git a/sqlspec/adapters/mysqlconnector/_typing.py b/sqlspec/adapters/mysqlconnector/_typing.py index 1ad16a74d..ad74c7fe8 100644 --- a/sqlspec/adapters/mysqlconnector/_typing.py +++ b/sqlspec/adapters/mysqlconnector/_typing.py @@ -17,8 +17,6 @@ from mysql.connector.constants import FieldType as _MysqlConnectorFieldType from mysql.connector.cursor import MySQLCursor as _MysqlConnectorSyncRawCursor -from sqlspec.typing import import_optional_attr - if TYPE_CHECKING: from collections.abc import Awaitable, Callable from types import TracebackType @@ -63,7 +61,7 @@ class MysqlConnectorMysqlModuleProtocol(Protocol): MysqlConnectorAsyncRawCursor: TypeAlias = _MysqlConnectorAsyncRawCursor if not TYPE_CHECKING: - MysqlConnectorAsyncPool = import_optional_attr("mysql.connector.aio.pooling", "MySQLConnectionPool") + MysqlConnectorAsyncPool = Any MysqlConnectorAio = _mysql_connector_aio MysqlConnectorSyncConnection = _MysqlConnectorSyncConnection MysqlConnectorAsyncConnection = _MysqlConnectorAsyncConnection diff --git a/sqlspec/adapters/mysqlconnector/adk/store.py b/sqlspec/adapters/mysqlconnector/adk/store.py index e490584e9..df8f23c7d 100644 --- a/sqlspec/adapters/mysqlconnector/adk/store.py +++ b/sqlspec/adapters/mysqlconnector/adk/store.py @@ -933,12 +933,12 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) else: @@ -949,11 +949,11 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) await cursor.execute(sql, params) @@ -1011,6 +1011,9 @@ async def search_entries( records: list[StoredMemory] = [] for row in rows: rec = cast("StoredMemory", dict(zip(columns, row, strict=False))) + rec["content_json"] = _json_dict(rec.get("content_json")) + metadata_val = rec.get("metadata_json") + rec["metadata_json"] = _json_dict(metadata_val) if metadata_val is not None else None rec["embedding"] = None records.append(rec) return records @@ -1146,12 +1149,12 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) else: @@ -1162,11 +1165,11 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) cursor.execute(sql, cast("tuple[Any, ...]", params)) @@ -1225,6 +1228,9 @@ def search_entries( records: list[StoredMemory] = [] for row in rows: rec = cast("StoredMemory", dict(zip(columns, row, strict=False))) + rec["content_json"] = _json_dict(rec.get("content_json")) + metadata_val = rec.get("metadata_json") + rec["metadata_json"] = _json_dict(metadata_val) if metadata_val is not None else None rec["embedding"] = None records.append(rec) return records diff --git a/sqlspec/adapters/mysqlconnector/config.py b/sqlspec/adapters/mysqlconnector/config.py index 58a0f8d85..903031934 100644 --- a/sqlspec/adapters/mysqlconnector/config.py +++ b/sqlspec/adapters/mysqlconnector/config.py @@ -397,7 +397,10 @@ def _create_pool(self) -> "MysqlConnectorConnectionPool": pool_size = config.pop("pool_size", None) pool_reset = config.pop("pool_reset_session", True) return MysqlConnectorConnectionPool( - pool_name=pool_name, pool_size=pool_size or 5, pool_reset_session=pool_reset, **config + pool_name=pool_name, + pool_size=5 if pool_size is None else pool_size, + pool_reset_session=pool_reset, + **config, ) def _close_pool(self) -> None: diff --git a/sqlspec/adapters/mysqlconnector/core.py b/sqlspec/adapters/mysqlconnector/core.py index d825bf913..8d92d0a3c 100644 --- a/sqlspec/adapters/mysqlconnector/core.py +++ b/sqlspec/adapters/mysqlconnector/core.py @@ -100,10 +100,14 @@ def __init__( def start(self) -> None: handler = self._driver.handle_database_exceptions() with handler: - cursor = self._driver.connection.cursor(**self._cursor_options) - self._cursor = cursor - cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + try: + cursor = self._driver.connection.cursor(**self._cursor_options) + self._cursor = cursor + cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) + self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + except BaseException: + self.close(error=True) + raise self._driver._check_pending_exception(handler) def fetch_chunk(self) -> "list[dict[str, Any]]": @@ -176,10 +180,14 @@ async def start(self) -> None: self._driver._check_pending_exception(handler) async def _start(self) -> None: - cursor = await self._driver.connection.cursor(**self._cursor_options) - self._cursor = cursor - await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + try: + cursor = await self._driver.connection.cursor(**self._cursor_options) + self._cursor = cursor + await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) + self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + except BaseException: + await self.close(error=True) + raise async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() diff --git a/sqlspec/adapters/oracledb/_json_handlers.py b/sqlspec/adapters/oracledb/_json_handlers.py index cbd9360b1..a8a959b5c 100644 --- a/sqlspec/adapters/oracledb/_json_handlers.py +++ b/sqlspec/adapters/oracledb/_json_handlers.py @@ -36,6 +36,7 @@ from functools import partial from typing import TYPE_CHECKING, Any +from sqlspec.adapters.oracledb._param_types import OracleJson from sqlspec.adapters.oracledb._typing import ( DB_TYPE_BLOB, DB_TYPE_CHAR, @@ -67,6 +68,7 @@ "is_json_payload", "json_converter_in_blob", "json_converter_in_clob", + "json_converter_in_native", "json_converter_out_blob", "json_converter_out_clob", "json_converter_out_oson", @@ -78,13 +80,24 @@ _JSON_STRING_TYPE_CODES = (DB_TYPE_VARCHAR, DB_TYPE_CHAR, DB_TYPE_NVARCHAR, DB_TYPE_NCHAR) +def json_converter_in_native(value: Any) -> Any: + """Unwrap an explicit ``OracleJson`` marker before ``DB_TYPE_JSON`` binding.""" + if isinstance(value, OracleJson): + return value.value + return value + + def json_converter_in_clob(value: Any) -> str: """Serialize a Python value to a JSON string for CLOB binding.""" + if isinstance(value, OracleJson): + value = value.value return to_json(value) def json_converter_in_blob(value: Any) -> bytes: """Serialize a Python value to UTF-8 JSON bytes for BLOB binding.""" + if isinstance(value, OracleJson): + value = value.value return to_json(value, as_bytes=True) @@ -172,14 +185,16 @@ def chain_output_handler(inner: Any, fallback: "Any | None") -> Any: def is_json_payload(value: Any) -> bool: """Return True if the value should be claimed by the JSON input handler. - ``dict`` and ``tuple``/``list`` of dicts are claimed. Sequences whose first - element is a number are NOT claimed — those are vector embeddings and - belong to the vector handler. + ``OracleJson``, ``dict``, and ``tuple``/``list`` of dicts are claimed. + Unwrapped sequences whose first element is a number are NOT claimed — those + are vector embeddings and belong to the vector handler. - An empty sequence is ambiguous (could be empty vector or empty list) and - defers to the next handler in the chain. Sequences of numbers (vector - embeddings) are rejected. + An unwrapped empty sequence is ambiguous (could be empty vector or empty + list) and defers to the next handler in the chain. Sequences of numbers + (vector embeddings) are rejected unless wrapped in ``OracleJson``. """ + if isinstance(value, OracleJson): + return True if isinstance(value, dict): return True if isinstance(value, (list, tuple)): @@ -198,10 +213,14 @@ def _input_type_handler(cursor: "Cursor | AsyncCursor", value: Any, arraysize: i server_major = resolve_oracle_connection_major(cursor.connection) if server_major is None: + if isinstance(value, OracleJson): + return cursor.var(DB_TYPE_JSON, arraysize=arraysize, inconverter=json_converter_in_native) return cursor.var(DB_TYPE_JSON, arraysize=arraysize) storage = resolve_oracle_json_storage(server_major) if storage == ORACLE_JSON_STORAGE_NATIVE: + if isinstance(value, OracleJson): + return cursor.var(DB_TYPE_JSON, arraysize=arraysize, inconverter=json_converter_in_native) return cursor.var(DB_TYPE_JSON, arraysize=arraysize) if storage == ORACLE_JSON_STORAGE_BLOB_JSON: return cursor.var(DB_TYPE_BLOB, arraysize=arraysize, inconverter=json_converter_in_blob) diff --git a/sqlspec/adapters/oracledb/adk/store.py b/sqlspec/adapters/oracledb/adk/store.py index c302b09d2..37b73dd43 100644 --- a/sqlspec/adapters/oracledb/adk/store.py +++ b/sqlspec/adapters/oracledb/adk/store.py @@ -172,8 +172,6 @@ " BEGIN\n" " EXECUTE IMMEDIATE 'CREATE TABLE {0} (\n" " id VARCHAR2(128) PRIMARY KEY,\n" - " app_name VARCHAR2(128) NOT NULL,\n" - " user_id VARCHAR2(128) NOT NULL,\n" " session_id VARCHAR2(128) NOT NULL,\n" " app_name VARCHAR2(128) NOT NULL,\n" " user_id VARCHAR2(128) NOT NULL,\n" @@ -408,9 +406,9 @@ async def create_session( params = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_data} async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql, params) + await conn.commit() result = await self.get_session(app_name, user_id, session_id) if result is None: @@ -439,39 +437,39 @@ async def get_session( try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + with conn.cursor() as cursor: + if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + await cursor.execute( + f"UPDATE {self._session_table} SET update_time = SYSTIMESTAMP WHERE app_name = :app_name AND user_id = :user_id AND id = :id", + {"app_name": app_name, "user_id": user_id, "id": session_id}, + ) + await conn.commit() + await cursor.execute( - f"UPDATE {self._session_table} SET update_time = SYSTIMESTAMP WHERE app_name = :app_name AND user_id = :user_id AND id = :id", + f""" + SELECT id, app_name, user_id, state, create_time, update_time + FROM {self._session_table} + WHERE app_name = :app_name AND user_id = :user_id AND id = :id + """, {"app_name": app_name, "user_id": user_id, "id": session_id}, ) - await conn.commit() + row = await cursor.fetchone() - await cursor.execute( - f""" - SELECT id, app_name, user_id, state, create_time, update_time - FROM {self._session_table} - WHERE app_name = :app_name AND user_id = :user_id AND id = :id - """, - {"app_name": app_name, "user_id": user_id, "id": session_id}, - ) - row = await cursor.fetchone() + if row is None: + return None - if row is None: - return None - - session_id_val, app_name, user_id, state_data, create_time, update_time = row + session_id_val, app_name, user_id, state_data, create_time, update_time = row - state = await self._deserialize_state(state_data) + state = await self._deserialize_state(state_data) - return StoredSession( - id=session_id_val, - app_name=app_name, - user_id=user_id, - state=state, - create_time=create_time, - update_time=update_time, - ) + return StoredSession( + id=session_id_val, + app_name=app_name, + user_id=user_id, + state=state, + create_time=create_time, + update_time=update_time, + ) except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -501,9 +499,11 @@ async def update_session_state(self, app_name: str, user_id: str, session_id: st """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute( + sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id} + ) + await conn.commit() async def list_sessions( self, @@ -542,25 +542,25 @@ async def list_sessions( try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - rows = await cursor.fetchall() - - results = [] - for row in rows: - state = await self._deserialize_state(row[3]) - - results.append( - StoredSession( - id=row[0], - app_name=row[1], - user_id=row[2], - state=state, - create_time=row[4], - update_time=row[5], + with conn.cursor() as cursor: + await cursor.execute(sql, params) + rows = await cursor.fetchall() + + results = [] + for row in rows: + state = await self._deserialize_state(row[3]) + + results.append( + StoredSession( + id=row[0], + app_name=row[1], + user_id=row[2], + state=state, + create_time=row[4], + update_time=row[5], + ) ) - ) - return results + return results except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -581,9 +581,9 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> sql = f"DELETE FROM {self._session_table} WHERE app_name = :app_name AND user_id = :user_id AND id = :id" async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"app_name": app_name, "user_id": user_id, "id": session_id}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql, {"app_name": app_name, "user_id": user_id, "id": session_id}) + await conn.commit() async def append_event(self, event_record: StoredEvent) -> None: """Append an event to a session. @@ -600,20 +600,20 @@ async def append_event(self, event_record: StoredEvent) -> None: """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute( - sql, - { - "id": event_record["id"], - "app_name": event_record["app_name"], - "user_id": event_record["user_id"], - "session_id": event_record["session_id"], - "invocation_id": event_record["invocation_id"], - "timestamp": event_record["timestamp"], - "event_data": await self._serialize_event_data(event_record["event_data"]), - }, - ) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute( + sql, + { + "id": event_record["id"], + "app_name": event_record["app_name"], + "user_id": event_record["user_id"], + "session_id": event_record["session_id"], + "invocation_id": event_record["invocation_id"], + "timestamp": event_record["timestamp"], + "event_data": await self._serialize_event_data(event_record["event_data"]), + }, + ) + await conn.commit() async def append_event_and_update_state( self, @@ -675,40 +675,44 @@ async def append_event_and_update_state( """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - try: - await cursor.execute( - update_sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id} - ) - await cursor.execute(select_sql, {"app_name": app_name, "user_id": user_id, "id": session_id}) - row = await cursor.fetchone() - if row is None: - _raise_session_not_found(session_id) - await cursor.execute( - insert_sql, - { - "id": event_record["id"], - "app_name": event_record["app_name"], - "user_id": event_record["user_id"], - "session_id": event_record["session_id"], - "invocation_id": event_record["invocation_id"], - "timestamp": event_record["timestamp"], - "event_data": await self._serialize_event_data(event_record["event_data"]), - }, - ) - if app_state is not None: + with conn.cursor() as cursor: + try: await cursor.execute( - app_upsert_sql, {"app_name": app_name, "state": await self._serialize_state(app_state)} + update_sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id} ) - if user_state is not None: + await cursor.execute(select_sql, {"app_name": app_name, "user_id": user_id, "id": session_id}) + row = await cursor.fetchone() + if row is None: + _raise_session_not_found(session_id) await cursor.execute( - user_upsert_sql, - {"app_name": app_name, "user_id": user_id, "state": await self._serialize_state(user_state)}, + insert_sql, + { + "id": event_record["id"], + "app_name": event_record["app_name"], + "user_id": event_record["user_id"], + "session_id": event_record["session_id"], + "invocation_id": event_record["invocation_id"], + "timestamp": event_record["timestamp"], + "event_data": await self._serialize_event_data(event_record["event_data"]), + }, ) - await conn.commit() - except Exception: - await conn.rollback() - raise + if app_state is not None: + await cursor.execute( + app_upsert_sql, {"app_name": app_name, "state": await self._serialize_state(app_state)} + ) + if user_state is not None: + await cursor.execute( + user_upsert_sql, + { + "app_name": app_name, + "user_id": user_id, + "state": await self._serialize_state(user_state), + }, + ) + await conn.commit() + except Exception: + await conn.rollback() + raise session_id_val, row_app_name, row_user_id, state_data_row, create_time, update_time = row return StoredSession( @@ -766,22 +770,22 @@ async def get_events( try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - rows = await cursor.fetchall() + with conn.cursor() as cursor: + await cursor.execute(sql, params) + rows = await cursor.fetchall() - return [ - StoredEvent( - id=row[0], - session_id=row[1], - invocation_id=_oracle_text_value(row[2]), - timestamp=row[3], - event_data=await self._deserialize_json_field(row[4]) or {}, - app_name=row[5], - user_id=row[6], - ) - for row in rows - ] + return [ + StoredEvent( + id=row[0], + session_id=row[1], + invocation_id=_oracle_text_value(row[2]), + timestamp=row[3], + event_data=await self._deserialize_json_field(row[4]) or {}, + app_name=row[5], + user_id=row[6], + ) + for row in rows + ] except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -797,10 +801,10 @@ async def delete_expired_events(self, before: "datetime", app_name: "str | None" try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - await conn.commit() - return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 + with conn.cursor() as cursor: + await cursor.execute(sql, params) + await conn.commit() + return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -816,10 +820,10 @@ async def delete_idle_sessions(self, updated_before: "datetime", app_name: "str try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - await conn.commit() - return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 + with conn.cursor() as cursor: + await cursor.execute(sql, params) + await conn.commit() + return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -835,10 +839,10 @@ async def delete_idle_user_states(self, updated_before: "datetime", app_name: "s try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - await conn.commit() - return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 + with conn.cursor() as cursor: + await cursor.execute(sql, params) + await conn.commit() + return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -851,10 +855,10 @@ async def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"app_name": app_name}) - row = await cursor.fetchone() - return await self._deserialize_state(row[0]) if row is not None else None + with conn.cursor() as cursor: + await cursor.execute(sql, {"app_name": app_name}) + row = await cursor.fetchone() + return await self._deserialize_state(row[0]) if row is not None else None except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -871,10 +875,10 @@ async def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"app_name": app_name, "user_id": user_id}) - row = await cursor.fetchone() - return await self._deserialize_state(row[0]) if row is not None else None + with conn.cursor() as cursor: + await cursor.execute(sql, {"app_name": app_name, "user_id": user_id}) + row = await cursor.fetchone() + return await self._deserialize_state(row[0]) if row is not None else None except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -895,9 +899,9 @@ async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"app_name": app_name, "state": await self._serialize_state(state)}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql, {"app_name": app_name, "state": await self._serialize_state(state)}) + await conn.commit() async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: """Insert or replace user-scoped state for an application user.""" @@ -913,11 +917,11 @@ async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute( - sql, {"app_name": app_name, "user_id": user_id, "state": await self._serialize_state(state)} - ) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute( + sql, {"app_name": app_name, "user_id": user_id, "state": await self._serialize_state(state)} + ) + await conn.commit() async def get_metadata(self, key: str) -> "str | None": """Return a value from the ADK internal metadata table.""" @@ -925,10 +929,10 @@ async def get_metadata(self, key: str) -> "str | None": try: async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"key": key}) - row = await cursor.fetchone() - return str(row[0]) if row is not None else None + with conn.cursor() as cursor: + await cursor.execute(sql, {"key": key}) + row = await cursor.fetchone() + return str(row[0]) if row is not None else None except OracleDatabaseError as e: error_obj = e.args[0] if e.args else None if error_obj and error_obj.code == ORACLE_TABLE_NOT_FOUND_ERROR: @@ -949,9 +953,9 @@ async def set_metadata(self, key: str, value: str) -> None: """ async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"key": key, "value": value}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql, {"key": key, "value": value}) + await conn.commit() async def _sessions_table_ddl(self) -> str: """Get Oracle CREATE TABLE SQL for sessions table. @@ -1411,8 +1415,7 @@ def create_session( """ params = {"id": session_id, "app_name": app_name, "user_id": user_id, "state": state_data} - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) conn.commit() @@ -1448,8 +1451,7 @@ def get_session( """ try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: if renew_for is not None and self._calculate_expires_at(renew_for) is not None: cursor.execute( f"UPDATE {self._session_table} SET update_time = SYSTIMESTAMP WHERE app_name = :app_name AND user_id = :user_id AND id = :id", @@ -1503,8 +1505,7 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta WHERE app_name = :app_name AND user_id = :user_id AND id = :id """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id}) conn.commit() @@ -1544,8 +1545,7 @@ def list_sessions( ) try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) rows = cursor.fetchall() @@ -1583,8 +1583,7 @@ def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: """ sql = f"DELETE FROM {self._session_table} WHERE app_name = :app_name AND user_id = :user_id AND id = :id" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"app_name": app_name, "user_id": user_id, "id": session_id}) conn.commit() @@ -1598,8 +1597,7 @@ def append_event(self, event_record: StoredEvent) -> None: ) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute( sql, { @@ -1669,8 +1667,7 @@ def append_event_and_update_state( VALUES (source.app_name, source.user_id, source.state, SYSTIMESTAMP) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: try: cursor.execute( update_sql, {"state": state_data, "app_name": app_name, "user_id": user_id, "id": session_id} @@ -1755,8 +1752,7 @@ def get_events( """ try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) rows = cursor.fetchall() @@ -1787,8 +1783,7 @@ def delete_expired_events(self, before: "datetime", app_name: "str | None" = Non params["app_name"] = app_name try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) conn.commit() return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 @@ -1807,8 +1802,7 @@ def delete_idle_sessions(self, updated_before: "datetime", app_name: "str | None params["app_name"] = app_name try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) conn.commit() return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 @@ -1827,8 +1821,7 @@ def delete_idle_user_states(self, updated_before: "datetime", app_name: "str | N params["app_name"] = app_name try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) conn.commit() return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 @@ -1843,8 +1836,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = :app_name" try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"app_name": app_name}) row = cursor.fetchone() return self._deserialize_state(row[0]) if row is not None else None @@ -1863,8 +1855,7 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" """ try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"app_name": app_name, "user_id": user_id}) row = cursor.fetchone() return self._deserialize_state(row[0]) if row is not None else None @@ -1887,8 +1878,7 @@ def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: VALUES (source.app_name, source.state, SYSTIMESTAMP) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"app_name": app_name, "state": self._serialize_state(state)}) conn.commit() @@ -1905,8 +1895,7 @@ def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]" VALUES (source.app_name, source.user_id, source.state, SYSTIMESTAMP) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"app_name": app_name, "user_id": user_id, "state": self._serialize_state(state)}) conn.commit() @@ -1915,8 +1904,7 @@ def get_metadata(self, key: str) -> "str | None": sql = f"SELECT value FROM {self._metadata_table} WHERE key = :key" try: - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"key": key}) row = cursor.fetchone() return str(row[0]) if row is not None else None @@ -1939,8 +1927,7 @@ def set_metadata(self, key: str, value: str) -> None: VALUES (source.key, source.value) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"key": key, "value": value}) conn.commit() @@ -2327,29 +2314,29 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " inserted_count = 0 async with self._config.provide_connection() as conn: - cursor = conn.cursor() - for entry in entries: - content_json = await self._serialize_json_field(entry["content_json"]) - metadata_json = await self._serialize_json_field(entry["metadata_json"]) - params = { - "id": entry["id"], - "session_id": entry["session_id"], - "app_name": entry["app_name"], - "user_id": entry["user_id"], - "scope": entry.get("scope", "user"), - "event_id": entry["event_id"], - "author": entry["author"], - "timestamp": entry["timestamp"], - "content_json": content_json, - "content_text": entry["content_text"], - "metadata_json": metadata_json, - "inserted_at": entry["inserted_at"], - } - if self._owner_id_column_name: - params["owner_id"] = str(owner_id) if owner_id is not None else None - if await self._execute_insert_entry(cursor, sql, params): - inserted_count += 1 - await conn.commit() + with conn.cursor() as cursor: + for entry in entries: + content_json = await self._serialize_json_field(entry["content_json"]) + metadata_json = await self._serialize_json_field(entry["metadata_json"]) + params = { + "id": entry["id"], + "session_id": entry["session_id"], + "app_name": entry["app_name"], + "user_id": entry["user_id"], + "scope": entry.get("scope", "user"), + "event_id": entry["event_id"], + "author": entry["author"], + "timestamp": entry["timestamp"], + "content_json": content_json, + "content_text": entry["content_text"], + "metadata_json": metadata_json, + "inserted_at": entry["inserted_at"], + } + if self._owner_id_column_name: + params["owner_id"] = str(owner_id) if owner_id is not None else None + if await self._execute_insert_entry(cursor, sql, params): + inserted_count += 1 + await conn.commit() return inserted_count @@ -2381,10 +2368,10 @@ async def search_entries( async def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = :session_id" async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": session_id}) - await conn.commit() - return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 + with conn.cursor() as cursor: + await cursor.execute(sql, {"session_id": session_id}) + await conn.commit() + return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 async def delete_entries_older_than( self, days: int, app_name: "str | None" = None, scope: "str | None" = None @@ -2401,10 +2388,10 @@ async def delete_entries_older_than( where_sql = " AND ".join(clauses) sql = f"DELETE FROM {self._memory_table} WHERE {where_sql}" async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - await conn.commit() - return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 + with conn.cursor() as cursor: + await cursor.execute(sql, params) + await conn.commit() + return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 async def _detect_json_storage_type(self) -> "JSONStorageType": return storage_type_from_version(await self._get_version_info()) @@ -2557,10 +2544,10 @@ async def _search_entries_fts( """ params = {**scope_params, "query": query, "limit": limit} async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - rows = await cursor.fetchall() - return await self._rows_to_records(rows) + with conn.cursor() as cursor: + await cursor.execute(sql, params) + rows = await cursor.fetchall() + return await self._rows_to_records(rows) async def _search_entries_simple( self, query: str, app_name: str, user_id: str, limit: int, scope_filter: Literal["all", "user", "app"] = "all" @@ -2582,10 +2569,10 @@ async def _search_entries_simple( pattern = f"%{query.lower()}%" params = {**scope_params, "pattern": pattern, "limit": limit} async with self._config.provide_connection() as conn: - cursor = conn.cursor() - await cursor.execute(sql, params) - rows = await cursor.fetchall() - return await self._rows_to_records(rows) + with conn.cursor() as cursor: + await cursor.execute(sql, params) + rows = await cursor.fetchall() + return await self._rows_to_records(rows) async def _rows_to_records(self, rows: "list[Any]") -> "list[StoredMemory]": records: list[StoredMemory] = [] @@ -2659,8 +2646,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object """ inserted_count = 0 - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: for entry in entries: content_json = self._serialize_json_field(entry["content_json"]) metadata_json = self._serialize_json_field(entry["metadata_json"]) @@ -2715,8 +2701,7 @@ def search_entries( def delete_entries_by_session(self, session_id: str) -> int: """Delete all memory entries for a specific session.""" sql = f"DELETE FROM {self._memory_table} WHERE session_id = :session_id" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"session_id": session_id}) conn.commit() return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 @@ -2734,8 +2719,7 @@ def delete_entries_older_than(self, days: int, app_name: "str | None" = None, sc where_sql = " AND ".join(clauses) sql = f"DELETE FROM {self._memory_table} WHERE {where_sql}" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) conn.commit() return cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 @@ -2890,11 +2874,10 @@ def _search_entries_fts( WHERE ROWNUM <= :limit """ params = {**scope_params, "query": query, "limit": limit} - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) rows = cursor.fetchall() - return self._rows_to_records(rows) + return self._rows_to_records(rows) def _search_entries_simple( self, query: str, app_name: str, user_id: str, limit: int, scope_filter: Literal["all", "user", "app"] = "all" @@ -2915,11 +2898,10 @@ def _search_entries_simple( """ pattern = f"%{query.lower()}%" params = {**scope_params, "pattern": pattern, "limit": limit} - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, params) rows = cursor.fetchall() - return self._rows_to_records(rows) + return self._rows_to_records(rows) def _rows_to_records(self, rows: "list[Any]") -> "list[StoredMemory]": records: list[StoredMemory] = [] diff --git a/sqlspec/adapters/oracledb/core.py b/sqlspec/adapters/oracledb/core.py index 44498f617..704e57c51 100644 --- a/sqlspec/adapters/oracledb/core.py +++ b/sqlspec/adapters/oracledb/core.py @@ -658,23 +658,27 @@ def __init__( def start(self) -> None: handler = self._driver.handle_database_exceptions() with handler: - cursor = self._driver.connection.cursor() - self._cursor = cursor - cursor.arraysize = self._chunk_size - cursor.prefetchrows = self._chunk_size - fetch_kwargs = build_fetch_kwargs(self._driver.driver_features) - if self._fetch_lobs is not None: - fetch_kwargs["fetch_lobs"] = self._fetch_lobs - parameters = coerce_large_parameters_sync( - self._driver.connection, - self._parameters, - clob_type=DB_TYPE_CLOB, - blob_type=DB_TYPE_BLOB, - varchar2_byte_limit=self._driver.driver_features.get("oracle_varchar2_byte_limit", 4000), - raw_byte_limit=self._driver.driver_features.get("oracle_raw_byte_limit", 2000), - version_cache=getattr(self._driver, "_oracle_version_cache", None), - ) - cast("Any", cursor).execute(self._sql, parameters or {}, **fetch_kwargs) + try: + cursor = self._driver.connection.cursor() + self._cursor = cursor + cursor.arraysize = self._chunk_size + cursor.prefetchrows = self._chunk_size + fetch_kwargs = build_fetch_kwargs(self._driver.driver_features) + if self._fetch_lobs is not None: + fetch_kwargs["fetch_lobs"] = self._fetch_lobs + parameters = coerce_large_parameters_sync( + self._driver.connection, + self._parameters, + clob_type=DB_TYPE_CLOB, + blob_type=DB_TYPE_BLOB, + varchar2_byte_limit=self._driver.driver_features.get("oracle_varchar2_byte_limit", 4000), + raw_byte_limit=self._driver.driver_features.get("oracle_raw_byte_limit", 2000), + version_cache=getattr(self._driver, "_oracle_version_cache", None), + ) + cast("Any", cursor).execute(self._sql, parameters or {}, **fetch_kwargs) + except BaseException: + self.close(error=True) + raise self._driver._check_pending_exception(handler) def fetch_chunk(self) -> "list[dict[str, Any]]": @@ -761,23 +765,27 @@ async def start(self) -> None: self._driver._check_pending_exception(handler) async def _start(self) -> None: - cursor = self._driver.connection.cursor() - self._cursor = cursor - cursor.arraysize = self._chunk_size - cursor.prefetchrows = self._chunk_size - fetch_kwargs = build_fetch_kwargs(self._driver.driver_features) - if self._fetch_lobs is not None: - fetch_kwargs["fetch_lobs"] = self._fetch_lobs - parameters = await coerce_large_parameters_async( - self._driver.connection, - self._parameters, - clob_type=DB_TYPE_CLOB, - blob_type=DB_TYPE_BLOB, - varchar2_byte_limit=self._driver.driver_features.get("oracle_varchar2_byte_limit", 4000), - raw_byte_limit=self._driver.driver_features.get("oracle_raw_byte_limit", 2000), - version_cache=getattr(self._driver, "_oracle_version_cache", None), - ) - await cast("Any", cursor).execute(self._sql, parameters or {}, **fetch_kwargs) + try: + cursor = self._driver.connection.cursor() + self._cursor = cursor + cursor.arraysize = self._chunk_size + cursor.prefetchrows = self._chunk_size + fetch_kwargs = build_fetch_kwargs(self._driver.driver_features) + if self._fetch_lobs is not None: + fetch_kwargs["fetch_lobs"] = self._fetch_lobs + parameters = await coerce_large_parameters_async( + self._driver.connection, + self._parameters, + clob_type=DB_TYPE_CLOB, + blob_type=DB_TYPE_BLOB, + varchar2_byte_limit=self._driver.driver_features.get("oracle_varchar2_byte_limit", 4000), + raw_byte_limit=self._driver.driver_features.get("oracle_raw_byte_limit", 2000), + version_cache=getattr(self._driver, "_oracle_version_cache", None), + ) + await cast("Any", cursor).execute(self._sql, parameters or {}, **fetch_kwargs) + except BaseException: + await self.close(error=True) + raise async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() @@ -1048,8 +1056,12 @@ def _parameter_values_need_coercion( if is_json_payload(value) and json_binding_state.uses_blob(): return True continue - if isinstance(value, (OracleClob, OracleBlob, OracleJson)): + if isinstance(value, (OracleClob, OracleBlob)): return True + if isinstance(value, OracleJson): + if json_binding_state.uses_blob(): + return True + continue if isinstance(value, str): if len(value.encode("utf-8")) > varchar2_byte_limit: return True @@ -1073,7 +1085,7 @@ def _coerce_parameters_sync( raw_byte_limit: int, json_binding_state: _OracleJsonBindingState, ) -> Any: - """Coerce one parameter container, copying sequences only when changed.""" + """Coerce one parameter container, copying mappings and sequences only when changed.""" if isinstance(parameters, dict): if not _parameter_values_need_coercion( parameters.values(), @@ -1082,6 +1094,7 @@ def _coerce_parameters_sync( json_binding_state=json_binding_state, ): return parameters + coerced_dict: dict[Any, Any] | None = None for param_name, param_value in parameters.items(): coerced_value = _coerce_value_sync( connection, @@ -1093,8 +1106,10 @@ def _coerce_parameters_sync( json_binding_state=json_binding_state, ) if coerced_value is not param_value: - parameters[param_name] = coerced_value - return parameters + if coerced_dict is None: + coerced_dict = dict(parameters) + coerced_dict[param_name] = coerced_value + return parameters if coerced_dict is None else coerced_dict if isinstance(parameters, (list, tuple)): if not _parameter_values_need_coercion( parameters, @@ -1142,6 +1157,7 @@ async def _coerce_parameters_async( json_binding_state=json_binding_state, ): return parameters + coerced_dict: dict[Any, Any] | None = None for param_name, param_value in parameters.items(): coerced_value = await _coerce_value_async( connection, @@ -1153,8 +1169,10 @@ async def _coerce_parameters_async( json_binding_state=json_binding_state, ) if coerced_value is not param_value: - parameters[param_name] = coerced_value - return parameters + if coerced_dict is None: + coerced_dict = dict(parameters) + coerced_dict[param_name] = coerced_value + return parameters if coerced_dict is None else coerced_dict if isinstance(parameters, (list, tuple)): if not _parameter_values_need_coercion( parameters, @@ -1224,15 +1242,9 @@ def _coerce_value_sync( inner = inner.encode("utf-8") return connection.createlob(blob_type, inner) if isinstance(value, OracleJson): - return _coerce_value_sync( - connection, - value.value, - clob_type=clob_type, - blob_type=blob_type, - varchar2_byte_limit=varchar2_byte_limit, - raw_byte_limit=raw_byte_limit, - json_binding_state=json_binding_state, - ) + if json_binding_state.uses_blob(): + return connection.createlob(blob_type, to_json(value.value, as_bytes=True)) + return value if isinstance(value, str) and len(value.encode("utf-8")) > varchar2_byte_limit: return connection.createlob(clob_type, value) if isinstance(value, (bytes, bytearray)) and len(value) > raw_byte_limit: @@ -1283,15 +1295,9 @@ async def _coerce_value_async( inner = inner.encode("utf-8") return await connection.createlob(blob_type, inner) if isinstance(value, OracleJson): - return await _coerce_value_async( - connection, - value.value, - clob_type=clob_type, - blob_type=blob_type, - varchar2_byte_limit=varchar2_byte_limit, - raw_byte_limit=raw_byte_limit, - json_binding_state=json_binding_state, - ) + if json_binding_state.uses_blob(): + return await connection.createlob(blob_type, to_json(value.value, as_bytes=True)) + return value if isinstance(value, str) and len(value.encode("utf-8")) > varchar2_byte_limit: return await connection.createlob(clob_type, value) if isinstance(value, (bytes, bytearray)) and len(value) > raw_byte_limit: diff --git a/sqlspec/adapters/oracledb/litestar/store.py b/sqlspec/adapters/oracledb/litestar/store.py index 7a62e9729..e0ee33c8b 100644 --- a/sqlspec/adapters/oracledb/litestar/store.py +++ b/sqlspec/adapters/oracledb/litestar/store.py @@ -145,31 +145,31 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": key}) - row = await cursor.fetchone() - - if row is None: - return None - - data_blob, expires_at = row - - if renew_for is not None and expires_at is not None: - expires_in_seconds = _oracle_expiry_seconds(renew_for) - if expires_in_seconds is not None: - update_sql = f""" - UPDATE {self._table_name} - SET expires_at = CASE - WHEN :expires_in_seconds IS NULL THEN NULL - ELSE SYSTIMESTAMP + NUMTODSINTERVAL(:expires_in_seconds, 'SECOND') - END, - updated_at = SYSTIMESTAMP - WHERE session_id = :session_id - """ - await cursor.execute(update_sql, {"expires_in_seconds": expires_in_seconds, "session_id": key}) - await conn.commit() - - return await _read_blob_async(data_blob) + with conn.cursor() as cursor: + await cursor.execute(sql, {"session_id": key}) + row = await cursor.fetchone() + + if row is None: + return None + + data_blob, expires_at = row + + if renew_for is not None and expires_at is not None: + expires_in_seconds = _oracle_expiry_seconds(renew_for) + if expires_in_seconds is not None: + update_sql = f""" + UPDATE {self._table_name} + SET expires_at = CASE + WHEN :expires_in_seconds IS NULL THEN NULL + ELSE SYSTIMESTAMP + NUMTODSINTERVAL(:expires_in_seconds, 'SECOND') + END, + updated_at = SYSTIMESTAMP + WHERE session_id = :session_id + """ + await cursor.execute(update_sql, {"expires_in_seconds": expires_in_seconds, "session_id": key}) + await conn.commit() + + return await _read_blob_async(data_blob) async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Store a session value. @@ -211,9 +211,11 @@ async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta conn_context = self._config.provide_connection() async with conn_context as conn: bind_data: Any = await conn.createlob(DB_TYPE_BLOB, data) if len(data) > ORACLE_SMALL_BLOB_LIMIT else data - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": key, "data": bind_data, "expires_in_seconds": expires_in_seconds}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute( + sql, {"session_id": key, "data": bind_data, "expires_in_seconds": expires_in_seconds} + ) + await conn.commit() async def delete(self, key: str) -> None: """Delete a session by key. @@ -225,9 +227,9 @@ async def delete(self, key: str) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": key}) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql, {"session_id": key}) + await conn.commit() async def delete_all(self) -> None: """Delete all sessions from the store.""" @@ -235,9 +237,9 @@ async def delete_all(self) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql) - await conn.commit() + with conn.cursor() as cursor: + await cursor.execute(sql) + await conn.commit() self._log_delete_all() async def exists(self, key: str) -> bool: @@ -257,10 +259,10 @@ async def exists(self, key: str) -> bool: conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": key}) - result = await cursor.fetchone() - return result is not None + with conn.cursor() as cursor: + await cursor.execute(sql, {"session_id": key}) + result = await cursor.fetchone() + return result is not None async def expires_in(self, key: str) -> "int | None": """Get the time in seconds until the session expires. @@ -278,25 +280,25 @@ async def expires_in(self, key: str) -> "int | None": conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql, {"session_id": key}) - row = await cursor.fetchone() + with conn.cursor() as cursor: + await cursor.execute(sql, {"session_id": key}) + row = await cursor.fetchone() - if row is None or row[0] is None: - return None + if row is None or row[0] is None: + return None - expires_at, db_now = row + expires_at, db_now = row - if expires_at.tzinfo is None: - expires_at = expires_at.replace(tzinfo=timezone.utc) - if db_now.tzinfo is None: - db_now = db_now.replace(tzinfo=timezone.utc) + if expires_at.tzinfo is None: + expires_at = expires_at.replace(tzinfo=timezone.utc) + if db_now.tzinfo is None: + db_now = db_now.replace(tzinfo=timezone.utc) - if expires_at <= db_now: - return 0 + if expires_at <= db_now: + return 0 - delta = expires_at - db_now - return int(delta.total_seconds()) + delta = expires_at - db_now + return int(delta.total_seconds()) async def delete_expired(self) -> int: """Delete all expired sessions. @@ -308,13 +310,13 @@ async def delete_expired(self) -> int: conn_context = self._config.provide_connection() async with conn_context as conn: - cursor = conn.cursor() - await cursor.execute(sql) - count = cursor.rowcount if cursor.rowcount is not None else 0 - await conn.commit() - if count > 0: - self._log_delete_expired(count) - return count + with conn.cursor() as cursor: + await cursor.execute(sql) + count = cursor.rowcount if cursor.rowcount is not None else 0 + await conn.commit() + if count > 0: + self._log_delete_expired(count) + return count def _table_ddl(self) -> str: """Get Oracle CREATE TABLE SQL with optimized schema. @@ -568,8 +570,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | AND (expires_at IS NULL OR expires_at > SYSTIMESTAMP) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"session_id": key}) row = cursor.fetchone() @@ -628,16 +629,15 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No with self._config.provide_connection() as conn: bind_data: Any = conn.createlob(DB_TYPE_BLOB, data) if len(data) > ORACLE_SMALL_BLOB_LIMIT else data - cursor = conn.cursor() - cursor.execute(sql, {"session_id": key, "data": bind_data, "expires_in_seconds": expires_in_seconds}) - conn.commit() + with conn.cursor() as cursor: + cursor.execute(sql, {"session_id": key, "data": bind_data, "expires_in_seconds": expires_in_seconds}) + conn.commit() def _delete(self, key: str) -> None: """Synchronous implementation of delete.""" sql = f"DELETE FROM {self._table_name} WHERE session_id = :session_id" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"session_id": key}) conn.commit() @@ -645,8 +645,7 @@ def _delete_all(self) -> None: """Synchronous implementation of delete_all.""" sql = f"DELETE FROM {self._table_name}" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql) conn.commit() self._log_delete_all() @@ -659,8 +658,7 @@ def _exists(self, key: str) -> bool: AND (expires_at IS NULL OR expires_at > SYSTIMESTAMP) """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"session_id": key}) result = cursor.fetchone() return result is not None @@ -672,8 +670,7 @@ def _expires_in(self, key: str) -> "int | None": WHERE session_id = :session_id """ - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql, {"session_id": key}) row = cursor.fetchone() @@ -697,8 +694,7 @@ def _delete_expired(self) -> int: """Synchronous implementation of delete_expired.""" sql = f"DELETE FROM {self._table_name} WHERE expires_at <= SYSTIMESTAMP" - with self._config.provide_connection() as conn: - cursor = conn.cursor() + with self._config.provide_connection() as conn, conn.cursor() as cursor: cursor.execute(sql) count = cursor.rowcount if cursor.rowcount is not None else 0 conn.commit() diff --git a/sqlspec/adapters/oracledb/type_converter.py b/sqlspec/adapters/oracledb/type_converter.py index 4d0b40c47..d2639a388 100644 --- a/sqlspec/adapters/oracledb/type_converter.py +++ b/sqlspec/adapters/oracledb/type_converter.py @@ -7,10 +7,19 @@ import array from typing import Any +from sqlspec.adapters.oracledb._vector_handlers import ( + numpy_converter_in, # pyright: ignore[reportPrivateUsage] + numpy_converter_out, # pyright: ignore[reportPrivateUsage] +) from sqlspec.typing import NUMPY_INSTALLED from sqlspec.utils.sync_tools import ensure_async_ from sqlspec.utils.type_guards import is_readable +try: + import numpy as np +except ImportError: + np = None # type: ignore[assignment] + __all__ = ("OracleOutputConverter",) @@ -54,8 +63,6 @@ def convert_vector_to_numpy(self, value: Any) -> Any: return value if isinstance(value, array.array): - from sqlspec.adapters.oracledb._vector_handlers import numpy_converter_out # pyright: ignore[reportPrivateUsage] - return numpy_converter_out(value) return value @@ -73,14 +80,10 @@ def convert_numpy_to_vector(self, value: Any) -> Any: array.array compatible with Oracle VECTOR if value is ndarray, otherwise original value. """ - if not NUMPY_INSTALLED: + if not NUMPY_INSTALLED or np is None: return value - import numpy as np - if isinstance(value, np.ndarray): - from sqlspec.adapters.oracledb._vector_handlers import numpy_converter_in # pyright: ignore[reportPrivateUsage] - return numpy_converter_in(value) return value diff --git a/sqlspec/adapters/psqlpy/_typing.py b/sqlspec/adapters/psqlpy/_typing.py index 3b04e90b0..8e2423cbb 100644 --- a/sqlspec/adapters/psqlpy/_typing.py +++ b/sqlspec/adapters/psqlpy/_typing.py @@ -44,8 +44,9 @@ class _PsqlpyUnavailableError(Exception): PsqlpyOperationalError: TypeAlias = _PsqlpyOperationalError -if not TYPE_CHECKING: +else: PsqlpyConnection = import_optional_attr("psqlpy", "Connection") or Any + PsqlpyConnectionPool = import_optional_attr("psqlpy", "ConnectionPool") or Any PsqlpyDataError = import_optional_attr("psqlpy.exceptions", "DataError") or _PsqlpyUnavailableError PsqlpyDatabaseError = import_optional_attr("psqlpy.exceptions", "DatabaseError") or _PsqlpyUnavailableError PsqlpyConnectionExecuteError = ( diff --git a/sqlspec/adapters/psqlpy/adk/store.py b/sqlspec/adapters/psqlpy/adk/store.py index d9503c6e9..970a4976c 100644 --- a/sqlspec/adapters/psqlpy/adk/store.py +++ b/sqlspec/adapters/psqlpy/adk/store.py @@ -663,21 +663,21 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " if self._owner_id_column_name: sql = f""" INSERT INTO {self._memory_table} ( - id, session_id, app_name, user_id, event_id, author, + id, session_id, app_name, user_id, scope, event_id, author, {self._owner_id_column_name}, timestamp, content_json, content_text, metadata_json, inserted_at ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12 + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13 ) ON CONFLICT (event_id) DO NOTHING """ else: sql = f""" INSERT INTO {self._memory_table} ( - id, session_id, app_name, user_id, event_id, author, + id, session_id, app_name, user_id, scope, event_id, author, timestamp, content_json, content_text, metadata_json, inserted_at ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11 + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12 ) ON CONFLICT (event_id) DO NOTHING """ @@ -690,6 +690,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["session_id"], entry["app_name"], entry["user_id"], + entry.get("scope", "user"), entry["event_id"], entry["author"], owner_id, @@ -705,6 +706,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["session_id"], entry["app_name"], entry["user_id"], + entry.get("scope", "user"), entry["event_id"], entry["author"], entry["timestamp"], @@ -735,15 +737,22 @@ async def search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) + if not query or not query.strip(): + return [] + effective_limit = limit if limit is not None else self._max_results try: if self._use_fts: try: - return await self._search_entries_fts(query, app_name, user_id, effective_limit) + return await self._search_entries_fts( + query, app_name, user_id, effective_limit, scope_filter=scope_filter + ) except Exception as exc: logger.warning("FTS search failed; falling back to simple search: %s", exc) - return await self._search_entries_simple(query, app_name, user_id, effective_limit) + return await self._search_entries_simple( + query, app_name, user_id, effective_limit, scope_filter=scope_filter + ) except Exception as e: if _is_psqlpy_database_error(e): error_msg = str(e).lower() @@ -774,21 +783,31 @@ async def delete_entries_older_than( self, days: int, app_name: "str | None" = None, scope: "str | None" = None ) -> int: """Delete memory entries older than specified days.""" + clauses = ["inserted_at < (CURRENT_TIMESTAMP - ($1::int * INTERVAL '1 day'))"] + params: list[Any] = [days] + if app_name is not None: + params.append(app_name) + clauses.append(f"app_name = ${len(params)}") + if scope is not None: + params.append(scope) + clauses.append(f"scope = ${len(params)}") + where_clause = " AND ".join(clauses) + count_sql = f""" SELECT COUNT(*) AS count FROM {self._memory_table} - WHERE inserted_at < CURRENT_TIMESTAMP - INTERVAL '{days} days' + WHERE {where_clause} """ delete_sql = f""" DELETE FROM {self._memory_table} - WHERE inserted_at < CURRENT_TIMESTAMP - INTERVAL '{days} days' + WHERE {where_clause} """ try: async with self._config.provide_connection() as conn: - count_result = await conn.fetch(count_sql, []) + count_result = await conn.fetch(count_sql, params) count_rows: list[dict[str, Any]] = count_result.result() if count_result else [] count = int(count_rows[0]["count"]) if count_rows else 0 - await conn.execute(delete_sql, []) + await conn.execute(delete_sql, params) return count except Exception as e: if _is_psqlpy_database_error(e): @@ -816,6 +835,7 @@ async def _memory_table_ddl(self) -> str: session_id VARCHAR(128) NOT NULL, app_name VARCHAR(128) NOT NULL, user_id VARCHAR(128) NOT NULL, + scope VARCHAR(16) NOT NULL DEFAULT 'user', event_id VARCHAR(128) NOT NULL UNIQUE, author VARCHAR(256){owner_id_line}, timestamp TIMESTAMPTZ NOT NULL, diff --git a/sqlspec/adapters/psqlpy/config.py b/sqlspec/adapters/psqlpy/config.py index dd99bc8e1..20dd41365 100644 --- a/sqlspec/adapters/psqlpy/config.py +++ b/sqlspec/adapters/psqlpy/config.py @@ -7,6 +7,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.psqlpy._typing import PsqlpyConnection, PsqlpyCursor, PsqlpySessionContext +from sqlspec.adapters.psqlpy._typing import PsqlpyConnectionPool as ConnectionPool from sqlspec.adapters.psqlpy.core import apply_driver_features, build_connection_config, default_statement_config from sqlspec.adapters.psqlpy.driver import PsqlpyDriver, PsqlpyExceptionHandler from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs @@ -25,7 +26,6 @@ from collections.abc import Awaitable, Callable from types import TracebackType - from sqlspec.adapters.psqlpy._typing import PsqlpyConnectionPool as ConnectionPool from sqlspec.core import StatementConfig __all__ = ("PsqlpyConfig", "PsqlpyConnectionParams", "PsqlpyCursor", "PsqlpyDriverFeatures", "PsqlpyPoolParams") @@ -304,8 +304,6 @@ async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: async def _create_pool(self) -> "ConnectionPool": """Create the actual async connection pool.""" - from sqlspec.adapters.psqlpy._typing import PsqlpyConnectionPool as ConnectionPool - return ConnectionPool(**build_connection_config(self.connection_config)) async def _close_pool(self) -> None: diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index 602482c92..84ba7b6e0 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -662,11 +662,14 @@ def _dml_count_query(sql: str) -> str | None: """Build a PostgreSQL query that returns an exact count for one DML statement.""" try: expression = sqlglot.parse_one(sql, dialect="postgres") - except ParseError as error: - msg = f"Unable to build psqlpy DML row count query: {error}" - raise SQLSpecError(msg) from error + except ParseError: + return None - if not isinstance(expression, (exp.Insert, exp.Update, exp.Delete)) or expression.args.get("returning"): + if ( + not isinstance(expression, (exp.Insert, exp.Update, exp.Delete)) + or expression.args.get("returning") + or expression.args.get("with_") + ): return None used_aliases = {cte.alias_or_name for cte in expression.find_all(exp.CTE)} diff --git a/sqlspec/adapters/psycopg/_typing.py b/sqlspec/adapters/psycopg/_typing.py index b30aa323a..9f3ffb2f8 100644 --- a/sqlspec/adapters/psycopg/_typing.py +++ b/sqlspec/adapters/psycopg/_typing.py @@ -51,8 +51,7 @@ PsycopgAsyncConnection: TypeAlias = AsyncConnection[PsycopgDictRow] PsycopgSyncRawCursor: TypeAlias = Cursor[PsycopgDictRow] PsycopgAsyncRawCursor: TypeAlias = AsyncCursor[PsycopgDictRow] - -if not TYPE_CHECKING: +else: PsycopgSyncConnection = Connection PsycopgAsyncConnection = AsyncConnection PsycopgSyncRawCursor = Cursor diff --git a/sqlspec/adapters/psycopg/adk/store.py b/sqlspec/adapters/psycopg/adk/store.py index 16ae88874..8267bbc75 100644 --- a/sqlspec/adapters/psycopg/adk/store.py +++ b/sqlspec/adapters/psycopg/adk/store.py @@ -225,6 +225,7 @@ async def create_tables(self) -> None: await driver.execute_script(await self._app_states_table_ddl()) await driver.execute_script(await self._user_states_table_ddl()) await driver.execute_script(await self._metadata_table_ddl()) + await driver.commit() async def create_session( self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None @@ -247,6 +248,7 @@ async def create_session( async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, params) + await conn.commit() result = await self.get_session(app_name, user_id, session_id) if result is None: @@ -257,7 +259,8 @@ async def create_session( async def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None ) -> "StoredSession | None": - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + should_touch = renew_for is not None and self._calculate_expires_at(renew_for) is not None + if should_touch: query = pg_sql.SQL(""" UPDATE {table} SET update_time = CURRENT_TIMESTAMP @@ -277,6 +280,8 @@ async def get_session( async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, params) row = await cur.fetchone() + if should_touch: + await conn.commit() if row is None: return None @@ -301,6 +306,7 @@ async def update_session_state(self, app_name: str, user_id: str, session_id: st async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, (Jsonb(state), app_name, user_id, session_id)) + await conn.commit() async def list_sessions( self, @@ -346,6 +352,7 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, (app_name, user_id, session_id)) + await conn.commit() async def append_event(self, event_record: StoredEvent) -> None: query = pg_sql.SQL(""" @@ -370,6 +377,7 @@ async def append_event(self, event_record: StoredEvent) -> None: jsonb_value, ), ) + await conn.commit() async def append_event_and_update_state( self, @@ -604,6 +612,7 @@ async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, (app_name, Jsonb(state))) + await conn.commit() async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: query = pg_sql.SQL(""" @@ -616,6 +625,7 @@ async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, (app_name, user_id, Jsonb(state))) + await conn.commit() async def get_metadata(self, key: str) -> "str | None": query = pg_sql.SQL("SELECT value FROM {table} WHERE key = %s").format( @@ -639,6 +649,7 @@ async def set_metadata(self, key: str, value: str) -> None: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute(query, (key, value)) + await conn.commit() async def _sessions_table_ddl(self) -> str: owner_id_line = "" @@ -728,6 +739,7 @@ def create_tables(self) -> None: driver.execute_script(self._app_states_table_ddl()) driver.execute_script(self._user_states_table_ddl()) driver.execute_script(self._metadata_table_ddl()) + driver.commit() def create_session( self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None @@ -751,6 +763,7 @@ def create_session( with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute(query, params) + conn.commit() result = self.get_session(app_name, user_id, session_id) if result is None: @@ -762,7 +775,8 @@ def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None ) -> "StoredSession | None": """Get session by ID.""" - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + should_touch = renew_for is not None and self._calculate_expires_at(renew_for) is not None + if should_touch: query = pg_sql.SQL(""" UPDATE {table} SET update_time = CURRENT_TIMESTAMP @@ -782,6 +796,8 @@ def get_session( with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute(query, params) row = cur.fetchone() + if should_touch: + conn.commit() if row is None: return None @@ -807,6 +823,7 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute(query, (Jsonb(state), app_name, user_id, session_id)) + conn.commit() def list_sessions( self, @@ -854,10 +871,10 @@ def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute(query, (app_name, user_id, session_id)) + conn.commit() def append_event(self, event_record: StoredEvent) -> None: """Append an event to a session.""" - """Synchronous implementation of append_event.""" self._insert_event(event_record) def append_event_and_update_state( @@ -1247,17 +1264,18 @@ def __init__(self, config: "PsycopgAsyncConfig") -> None: async def create_tables(self) -> None: """Create the memory table and indexes if they don't exist.""" - if not self.create_schema_enabled: - await self.reconcile_schema() + if not self._enabled: return - if not self._enabled: + if not self.create_schema_enabled: + await self.reconcile_schema() return async with self._config.provide_session() as driver: if self._enable_bm25: self._config._ensure_pg_textsearch_available() await driver.execute_script(await self._memory_table_ddl()) + await driver.commit() async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: """Bulk insert memory entries with deduplication.""" @@ -1301,6 +1319,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " await cur.execute(query, _build_insert_params(entry)) if cur.rowcount and cur.rowcount > 0: inserted_count += cur.rowcount + await conn.commit() return inserted_count @@ -1318,13 +1337,14 @@ async def search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) - if not query and embedding is None: + has_query = bool(query and query.strip()) + if not has_query and embedding is None: return [] effective_limit = limit if limit is not None else self._max_results try: - if embedding is not None and self._enable_bm25 and query: + if embedding is not None and self._enable_bm25 and has_query: return await self._search_entries_hybrid( query, app_name, user_id, effective_limit, scope_filter, embedding ) @@ -1341,18 +1361,27 @@ async def search_entries( async def delete_entries_by_session(self, session_id: str) -> int: """Delete all memory entries for a specific session.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + sql = pg_sql.SQL("DELETE FROM {table} WHERE session_id = %s").format( table=pg_sql.Identifier(self._memory_table) ) async with self._config.provide_connection() as conn, conn.cursor() as cur: await cur.execute(sql, (session_id,)) + await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 async def delete_entries_older_than( self, days: int, app_name: "str | None" = None, scope: "str | None" = None ) -> int: """Delete memory entries older than specified days.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + clauses: list[pg_sql.Composable] = [ pg_sql.SQL("inserted_at < CURRENT_TIMESTAMP - {interval}::interval").format( interval=pg_sql.Literal(f"{days} days") @@ -1373,6 +1402,7 @@ async def delete_entries_older_than( async with self._config.provide_connection() as conn, conn.cursor() as cur: await cur.execute(sql, tuple(params) if params else None) + await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 async def _memory_table_ddl(self) -> str: @@ -1514,17 +1544,18 @@ class PsycopgSyncADKMemoryStore(BaseSyncADKMemoryStore["PsycopgSyncConfig"]): def create_tables(self) -> None: """Create the memory table and indexes if they don't exist.""" - if not self.create_schema_enabled: - self.reconcile_schema() + if not self._enabled: return - if not self._enabled: + if not self.create_schema_enabled: + self.reconcile_schema() return with self._config.provide_session() as driver: if self._enable_bm25: self._config._ensure_pg_textsearch_available() driver.execute_script(self._memory_table_ddl()) + driver.commit() def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: """Bulk insert memory entries with deduplication.""" @@ -1568,6 +1599,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object cur.execute(query, _build_insert_params(entry)) if cur.rowcount and cur.rowcount > 0: inserted_count += cur.rowcount + conn.commit() return inserted_count @@ -1585,13 +1617,14 @@ def search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) - if not query and embedding is None: + has_query = bool(query and query.strip()) + if not has_query and embedding is None: return [] effective_limit = limit if limit is not None else self._max_results try: - if embedding is not None and self._enable_bm25 and query: + if embedding is not None and self._enable_bm25 and has_query: return self._search_entries_hybrid(query, app_name, user_id, effective_limit, scope_filter, embedding) if embedding is not None: return self._search_entries_vector(app_name, user_id, effective_limit, scope_filter, embedding) @@ -1606,17 +1639,25 @@ def search_entries( def delete_entries_by_session(self, session_id: str) -> int: """Delete all memory entries for a specific session.""" - """Delete all memory entries for a specific session.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + sql = pg_sql.SQL("DELETE FROM {table} WHERE session_id = %s").format( table=pg_sql.Identifier(self._memory_table) ) with self._config.provide_connection() as conn, conn.cursor() as cur: cur.execute(sql, (session_id,)) + conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: """Delete memory entries older than specified days.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + clauses: list[pg_sql.Composable] = [ pg_sql.SQL("inserted_at < CURRENT_TIMESTAMP - {interval}::interval").format( interval=pg_sql.Literal(f"{days} days") @@ -1637,6 +1678,7 @@ def delete_entries_older_than(self, days: int, app_name: "str | None" = None, sc with self._config.provide_connection() as conn, conn.cursor() as cur: cur.execute(sql, tuple(params) if params else None) + conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 def _memory_table_ddl(self) -> str: @@ -1779,12 +1821,12 @@ def _build_insert_params(entry: "StoredMemory") -> "tuple[object, ...]": entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], entry.get("embedding"), Jsonb(entry["content_json"]), entry["content_text"], - Jsonb(entry["metadata_json"]) if entry["metadata_json"] is not None else None, + Jsonb(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) @@ -1797,13 +1839,13 @@ def _build_insert_params_with_owner(entry: "StoredMemory", owner_id: "object | N entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], entry.get("embedding"), Jsonb(entry["content_json"]), entry["content_text"], - Jsonb(entry["metadata_json"]) if entry["metadata_json"] is not None else None, + Jsonb(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index 660314a61..a4d9923ce 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -8,6 +8,7 @@ from psycopg.types.json import set_json_dumps, set_json_loads from typing_extensions import NotRequired, Self +import sqlspec.adapters.psycopg._typing as _psycopg_typing from sqlspec.adapters.psycopg._typing import ( PsycopgAsyncConnection, PsycopgAsyncCursor, @@ -318,7 +319,6 @@ def __init__( statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, driver_features) - # Extract user connection hook before storing driver_features features_dict = dict(driver_features) if driver_features else {} self._user_connection_hook: Callable[[PsycopgSyncConnection], None] | None = features_dict.pop( "on_connection_create", None @@ -372,10 +372,8 @@ def _setup_alloydb_connector( self, config: "dict[str, Any]", pool_parameters: "dict[str, Any] | None" = None ) -> None: """Setup AlloyDB connector and configure psycopg-pool connection_class.""" - from sqlspec.adapters.psycopg._typing import PsycopgAlloydbConnector as Connector - if self._alloydb_connector is None: - self._alloydb_connector = Connector() + self._alloydb_connector = _psycopg_typing.PsycopgAlloydbConnector() user = config.get("user") password = config.get("password") @@ -476,11 +474,9 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if self._pgvector_available: register_pgvector_sync(conn) - # Ensure connection is not left in INTRANS state from extension detection or registration if not conn.autocommit: conn.rollback() - # Call user-provided callback after internal setup if self._user_connection_hook is not None: self._user_connection_hook(conn) @@ -619,7 +615,6 @@ def __init__(self, config: "PsycopgAsyncConfig") -> None: async def __aenter__(self) -> "PsycopgAsyncConnection": if self._config.connection_instance is None: self._config.connection_instance = await self._config.create_pool() - # pool.connection() returns an async context manager if self._config.connection_instance: self._ctx = self._config.connection_instance.connection() return cast("PsycopgAsyncConnection", await self._ctx.__aenter__()) @@ -706,7 +701,6 @@ def __init__( statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, driver_features) - # Extract user connection hook before storing driver_features features_dict = dict(driver_features) if driver_features else {} self._user_connection_hook: Callable[[PsycopgAsyncConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None @@ -807,11 +801,9 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if self._pgvector_available: await register_pgvector_async(conn) - # Ensure connection is not left in INTRANS state from extension detection or registration if not conn.autocommit: await conn.rollback() - # Call user-provided callback after internal setup if self._user_connection_hook is not None: await self._user_connection_hook(conn) diff --git a/sqlspec/adapters/psycopg/type_converter.py b/sqlspec/adapters/psycopg/type_converter.py index 9b37165e3..5870ef344 100644 --- a/sqlspec/adapters/psycopg/type_converter.py +++ b/sqlspec/adapters/psycopg/type_converter.py @@ -9,6 +9,8 @@ from typing import TYPE_CHECKING, Any +from sqlspec.adapters.psycopg._typing import PsycopgProgrammingError as ProgrammingError +from sqlspec.adapters.psycopg._typing import psycopg_errors as errors from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import import_optional @@ -33,8 +35,6 @@ def register_pgvector_sync(connection: "Connection[Any]") -> None: Args: connection: Psycopg sync connection. """ - from sqlspec.adapters.psycopg._typing import PsycopgProgrammingError as ProgrammingError - pgvector_psycopg = _pgvector_psycopg if pgvector_psycopg is None: return @@ -58,8 +58,6 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: Args: connection: Psycopg async connection. """ - from sqlspec.adapters.psycopg._typing import PsycopgProgrammingError as ProgrammingError - pgvector_psycopg = _pgvector_psycopg if pgvector_psycopg is None: return @@ -84,8 +82,6 @@ def _is_missing_vector_error(error: Exception) -> bool: Returns: True if error indicates vector type not found. """ - from sqlspec.adapters.psycopg._typing import psycopg_errors as errors - message = str(error).lower() return ( "vector type not found" in message diff --git a/sqlspec/adapters/pymssql/adk/store.py b/sqlspec/adapters/pymssql/adk/store.py index 7d1737f93..74d9593b5 100644 --- a/sqlspec/adapters/pymssql/adk/store.py +++ b/sqlspec/adapters/pymssql/adk/store.py @@ -59,6 +59,11 @@ def create_tables(self) -> None: return with self._config.provide_session() as driver: + if self._json_column_type is None: + configured = _configured_json_column_type(self._native_json) + self._json_column_type = ( + configured if configured is not None else _json_column_type_from_sync_driver(driver) + ) driver.execute_script(self._sessions_table_ddl()) driver.execute_script(self._events_table_ddl()) driver.execute_script(self._app_states_table_ddl()) @@ -390,8 +395,11 @@ def _json_column_type_sync(self) -> str: if configured is not None: self._json_column_type = configured return configured - with self._config.provide_session() as driver: - self._json_column_type = _json_column_type_from_sync_driver(driver) + try: + with self._config.provide_session() as driver: + self._json_column_type = _json_column_type_from_sync_driver(driver) + except Exception: + return JSON_FALLBACK_COLUMN_TYPE return self._json_column_type def _execute_fetchone(self, sql: str, params: "tuple[Any, ...]" = (), *, commit: bool = False) -> "Any | None": @@ -452,7 +460,6 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object owner_column = f", {_quote_identifier(self._owner_id_column_name)}" if self._owner_id_column_name else "" owner_value = ", %s" if self._owner_id_column_name else "" - # Keep the key-range lock and insertion in one statement, including autocommit. sql = f""" INSERT INTO {_table_ref(self._memory_table)} ( id, session_id, app_name, user_id, scope, event_id, author, timestamp, @@ -478,7 +485,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry.get("metadata_json")), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, ) if self._owner_id_column_name: params = (*params, owner_id) @@ -601,6 +608,8 @@ def _adk_config(config: Any) -> PymssqlADKConfig: def _configured_json_column_type(native_json: "bool | None") -> "str | None": + if native_json is None: + return None if native_json is True: return JSON_NATIVE_COLUMN_TYPE return JSON_FALLBACK_COLUMN_TYPE diff --git a/sqlspec/adapters/pymssql/core.py b/sqlspec/adapters/pymssql/core.py index 6be2f5796..bccd94c81 100644 --- a/sqlspec/adapters/pymssql/core.py +++ b/sqlspec/adapters/pymssql/core.py @@ -278,7 +278,7 @@ def extract_error_number(exc: BaseException | None) -> int | None: return val if exc.args: first = exc.args[0] - if isinstance(first, int) and not isinstance(first, bool): + if isinstance(first, int) and not isinstance(first, bool) and first != 0: return first matches = _ERROR_NUMBER_PATTERN.findall(str(exc)) if not matches: diff --git a/sqlspec/adapters/pymssql/driver.py b/sqlspec/adapters/pymssql/driver.py index 0655b8052..ccf488b90 100644 --- a/sqlspec/adapters/pymssql/driver.py +++ b/sqlspec/adapters/pymssql/driver.py @@ -222,11 +222,12 @@ def rollback(self) -> None: cursor.execute("IF @@TRANCOUNT > 0 ROLLBACK TRANSACTION") else: self.connection.rollback() - self._explicit_transaction = False - self._transaction_active = False except PymssqlError as exc: msg = f"Failed to rollback SQL Server transaction: {exc}" raise SQLSpecError(msg) from exc + finally: + self._explicit_transaction = False + self._transaction_active = False def with_cursor(self, connection: "PymssqlConnection") -> "PymssqlCursor": return PymssqlCursor(connection) diff --git a/sqlspec/adapters/pymssql/pool.py b/sqlspec/adapters/pymssql/pool.py index a89b17e46..8ffabcc70 100644 --- a/sqlspec/adapters/pymssql/pool.py +++ b/sqlspec/adapters/pymssql/pool.py @@ -204,12 +204,8 @@ def release(self, connection: PymssqlConnection) -> None: _ = connection def size(self) -> int: - try: - _ = self._thread_local.connection - except AttributeError: - return 0 - else: - return 1 + with self._registry_lock: + return len(self._connection_registry) def checked_out(self) -> int: return 0 diff --git a/sqlspec/adapters/pymysql/_typing.py b/sqlspec/adapters/pymysql/_typing.py index e3debdc90..d3b4463ce 100644 --- a/sqlspec/adapters/pymysql/_typing.py +++ b/sqlspec/adapters/pymysql/_typing.py @@ -8,9 +8,12 @@ from typing import TYPE_CHECKING, Any import pymysql +from pymysql import MySQLError as PyMysqlMySQLError +from pymysql.connections import Connection as PyMysqlConnection from pymysql.constants import FIELD_TYPE as _PYMYSQL_FIELD_TYPE from pymysql.constants import SERVER_STATUS as _PYMYSQL_SERVER_STATUS from pymysql.cursors import RE_INSERT_VALUES as PYMYSQL_INSERT_VALUES_PATTERN +from pymysql.cursors import Cursor as PyMysqlRawCursor from pymysql.cursors import DictCursor as PyMysqlDictCursor from pymysql.cursors import SSCursor as PyMysqlSSCursor @@ -33,20 +36,23 @@ class PyMysqlServerStatusProtocol(Protocol): SERVER_STATUS_IN_TRANS: int PyMysqlConnect: TypeAlias = type["PyMysqlConnection"] - PyMysqlConnection: TypeAlias = pymysql.connections.Connection PyMysqlFieldType: TypeAlias = PyMysqlFieldTypeProtocol - PyMysqlMySQLError: TypeAlias = pymysql.MySQLError - PyMysqlRawCursor: TypeAlias = pymysql.cursors.Cursor PyMysqlServerStatus: TypeAlias = PyMysqlServerStatusProtocol if not TYPE_CHECKING: PyMysqlConnect = pymysql.connect - PyMysqlConnection = pymysql.connections.Connection PyMysqlFieldType = _PYMYSQL_FIELD_TYPE - PyMysqlMySQLError = pymysql.MySQLError - PyMysqlRawCursor = pymysql.cursors.Cursor PyMysqlServerStatus = _PYMYSQL_SERVER_STATUS + def _pymysql_cloud_sql_connector(*args: Any, **kwargs: Any) -> Any: + connector_cls = import_optional_attr("google.cloud.sql.connector", "Connector") + if connector_cls is None: + msg = "Cannot import 'Connector' from 'google.cloud.sql.connector'" + raise ImportError(msg) + return connector_cls(*args, **kwargs) + + PyMysqlCloudSqlConnector = _pymysql_cloud_sql_connector + __all__ = ( "PYMYSQL_INSERT_VALUES_PATTERN", @@ -128,22 +134,3 @@ def __exit__( self._release_connection(self._connection, exc_type=exc_type, exc_val=exc_val, exc_tb=exc_tb) self._connection = None return None - - -_LAZY_DRIVER_EXPORTS: dict[str, tuple[str, str]] = { - "PyMysqlCloudSqlConnector": ("google.cloud.sql.connector", "Connector") -} - - -def __getattr__(name: str) -> Any: - """Resolve optional driver symbols only when a consumer requests them.""" - target = _LAZY_DRIVER_EXPORTS.get(name) - if target is None: - msg = f"module {__name__!r} has no attribute {name!r}" - raise AttributeError(msg) - module_name, attribute = target - value = import_optional_attr(module_name, attribute) - if value is None: - msg = f"Cannot import {attribute!r} from {module_name!r}" - raise ImportError(msg) - return value diff --git a/sqlspec/adapters/pymysql/adk/store.py b/sqlspec/adapters/pymysql/adk/store.py index ba637540c..a9eb19214 100644 --- a/sqlspec/adapters/pymysql/adk/store.py +++ b/sqlspec/adapters/pymysql/adk/store.py @@ -343,12 +343,12 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), owner_id, entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) else: @@ -359,11 +359,11 @@ def _insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "objec entry["user_id"], entry.get("scope", "user"), entry["event_id"], - entry["author"], + entry.get("author"), entry["timestamp"], to_json(entry["content_json"]), entry["content_text"], - to_json(entry["metadata_json"]), + to_json(entry["metadata_json"]) if entry.get("metadata_json") is not None else None, entry["inserted_at"], ) cursor.execute(sql, params) @@ -420,6 +420,9 @@ def _search_entries( records: list[StoredMemory] = [] for row in rows: record = cast("StoredMemory", dict(zip(columns, row, strict=False))) + record["content_json"] = _json_dict(record.get("content_json")) + metadata_val = record.get("metadata_json") + record["metadata_json"] = _json_dict(metadata_val) if metadata_val is not None else None record["embedding"] = None records.append(record) return records diff --git a/sqlspec/adapters/pymysql/config.py b/sqlspec/adapters/pymysql/config.py index 17b60732d..f4a01fdc0 100644 --- a/sqlspec/adapters/pymysql/config.py +++ b/sqlspec/adapters/pymysql/config.py @@ -6,6 +6,7 @@ from typing_extensions import NotRequired +from sqlspec.adapters.pymysql._typing import PyMysqlCloudSqlConnector as Connector from sqlspec.adapters.pymysql._typing import PyMysqlConnection, PyMysqlCursor, PyMysqlRawCursor, PyMysqlSessionContext from sqlspec.adapters.pymysql.core import apply_driver_features, default_statement_config from sqlspec.adapters.pymysql.driver import PyMysqlDriver, PyMysqlExceptionHandler @@ -214,7 +215,7 @@ def __init__( self._driver_kwargs = driver_kwargs def __call__(self) -> "PyMysqlConnection": - connector = self._config.get_cloud_sql_connector() + connector = self._config._get_cloud_sql_connector() if connector is None: msg = "Cloud SQL connector is not initialized" raise ImproperConfigurationError(msg) @@ -308,7 +309,7 @@ def __init__( self._cloud_sql_connector: Any | None = None self._validate_connector_config() - def get_cloud_sql_connector(self) -> Any | None: + def _get_cloud_sql_connector(self) -> Any | None: """Return the configured Cloud SQL connector instance.""" return self._cloud_sql_connector @@ -332,8 +333,6 @@ def _validate_connector_config(self) -> None: def _setup_cloud_sql_connector(self, config: "dict[str, Any]") -> "_PyMysqlCloudSqlConnector": """Setup Cloud SQL connector and return a pool connection factory.""" - from sqlspec.adapters.pymysql._typing import PyMysqlCloudSqlConnector as Connector - self._cloud_sql_connector = Connector() user = config.get("user") diff --git a/sqlspec/adapters/pymysql/core.py b/sqlspec/adapters/pymysql/core.py index b8b9350c0..98de2c636 100644 --- a/sqlspec/adapters/pymysql/core.py +++ b/sqlspec/adapters/pymysql/core.py @@ -81,10 +81,14 @@ def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int, json def start(self) -> None: handler = self._driver.handle_database_exceptions() with handler: - cursor = self._driver.connection.cursor(PyMysqlSSCursor) - self._cursor = cursor - cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + try: + cursor = self._driver.connection.cursor(PyMysqlSSCursor) + self._cursor = cursor + cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) + self._row_plan = resolve_row_plan(self._cursor.description, self._json_type_codes) + except BaseException: + self.close(error=True) + raise self._driver._check_pending_exception(handler) def fetch_chunk(self) -> "list[dict[str, Any]]": diff --git a/sqlspec/adapters/pymysql/litestar/store.py b/sqlspec/adapters/pymysql/litestar/store.py index a33de33cb..7ef6b120c 100644 --- a/sqlspec/adapters/pymysql/litestar/store.py +++ b/sqlspec/adapters/pymysql/litestar/store.py @@ -5,6 +5,7 @@ from typing_extensions import NotRequired +from sqlspec.adapters.pymysql._typing import PyMysqlDictCursor, PyMysqlMySQLError from sqlspec.config import LitestarConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.litestar.store import BaseSQLSpecStore @@ -101,8 +102,6 @@ def _create_table(self) -> None: self._log_table_created() def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": - from sqlspec.adapters.pymysql._typing import PyMysqlDictCursor, PyMysqlMySQLError - sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = %s @@ -177,8 +176,6 @@ def _delete(self, key: str) -> None: conn.commit() def _delete_all(self) -> None: - from sqlspec.adapters.pymysql._typing import PyMysqlMySQLError - sql = f"DELETE FROM {self._table_name}" try: @@ -197,8 +194,6 @@ def _delete_all(self) -> None: raise def _exists(self, key: str) -> bool: - from sqlspec.adapters.pymysql._typing import PyMysqlMySQLError - sql = f""" SELECT 1 FROM {self._table_name} WHERE session_id = %s diff --git a/sqlspec/adapters/spanner/adk/store.py b/sqlspec/adapters/spanner/adk/store.py index 5facfb347..c2862f18d 100644 --- a/sqlspec/adapters/spanner/adk/store.py +++ b/sqlspec/adapters/spanner/adk/store.py @@ -2,10 +2,13 @@ from collections.abc import Iterable, Mapping from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Protocol, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, Protocol, cast +import sqlglot +from sqlglot import exp from typing_extensions import NotRequired, TypedDict +import sqlspec.dialects.spanner # noqa: F401 from sqlspec.adapters.spanner._typing import SpannerNotFound as NotFound from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.adapters.spanner.config import SpannerSyncConfig @@ -26,8 +29,6 @@ __all__ = ("SpannerADKConfig", "SpannerADKRetentionConfig", "SpannerSyncADKMemoryStore", "SpannerSyncADKStore") SPANNER_PARAM_TYPES: SpannerParamTypesProtocol = cast("SpannerParamTypesProtocol", param_types) -MIN_DROP_TABLE_TOKENS: Final = 3 -MIN_DROP_SEARCH_INDEX_TOKENS: Final = 4 class SpannerADKRetentionConfig(TypedDict): @@ -1185,27 +1186,29 @@ def _filter_existing_spanner_drops(statements: "list[str]", existing_tables: "se def _spanner_drop_statement_table(statement: str, existing_tables: "set[str]") -> "str | None": - tokens = statement.strip().split() - if len(tokens) >= MIN_DROP_TABLE_TOKENS and tokens[0].upper() == "DROP" and tokens[1].upper() == "TABLE": - table_name = tokens[2] - return table_name if table_name in existing_tables else None - - index_name: str | None = None - if len(tokens) >= MIN_DROP_TABLE_TOKENS and tokens[0].upper() == "DROP" and tokens[1].upper() == "INDEX": - index_name = tokens[2] - if ( - len(tokens) >= MIN_DROP_SEARCH_INDEX_TOKENS - and tokens[0].upper() == "DROP" - and tokens[1].upper() == "SEARCH" - and tokens[2].upper() == "INDEX" - ): - index_name = tokens[3] - if index_name is None: + try: + parsed = sqlglot.parse_one(statement, read="spanner") + if isinstance(parsed, exp.Command) and str(parsed.this).upper() == "DROP": + expr_sql = str(parsed.expression or "").strip() + if expr_sql.upper().startswith("SEARCH "): + parsed = sqlglot.parse_one(f"DROP {expr_sql[7:]}", read="spanner") + except Exception: + return None + + if not isinstance(parsed, exp.Drop): + return None + + target = parsed.this if isinstance(parsed.this, exp.Table) else parsed.find(exp.Table) + if target is None or not target.name: return None - for table_name in existing_tables: - if index_name.startswith(f"idx_{table_name}_"): - return table_name + kind = str(parsed.args.get("kind") or "").upper() + if kind == "TABLE": + return target.name if target.name in existing_tables else None + if kind in {"INDEX", "SEARCH INDEX"}: + for table_name in existing_tables: + if target.name.startswith(f"idx_{table_name}_"): + return table_name return None diff --git a/sqlspec/adapters/spanner/config.py b/sqlspec/adapters/spanner/config.py index c4b5f9328..b32286a87 100644 --- a/sqlspec/adapters/spanner/config.py +++ b/sqlspec/adapters/spanner/config.py @@ -5,7 +5,11 @@ from typing_extensions import NotRequired +from sqlspec.adapters.spanner._typing import SpannerBurstyPool as BurstyPool +from sqlspec.adapters.spanner._typing import SpannerClient as Client from sqlspec.adapters.spanner._typing import SpannerConnection +from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool +from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool from sqlspec.adapters.spanner._typing import SpannerTransactionType as TransactionType from sqlspec.adapters.spanner.core import apply_driver_features, default_statement_config from sqlspec.adapters.spanner.driver import SpannerSessionContext, SpannerSyncDriver @@ -23,7 +27,6 @@ from types import TracebackType from sqlspec.adapters.spanner._typing import SpannerAbstractSessionPool as AbstractSessionPool - from sqlspec.adapters.spanner._typing import SpannerClient as Client from sqlspec.adapters.spanner._typing import SpannerClientInfo as ClientInfo from sqlspec.adapters.spanner._typing import SpannerClientOptions as ClientOptions from sqlspec.adapters.spanner._typing import SpannerCredentials as Credentials @@ -325,8 +328,6 @@ def __init__( ): self.connection_config["session_labels"] = legacy_session_labels - from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool - self.connection_config.setdefault("size", self.connection_config.pop("max_sessions", 10)) self.connection_config.setdefault("pool_type", FixedSizePool) @@ -349,8 +350,6 @@ def __init__( self._database: Database | None = None def _get_client(self) -> "Client": - from sqlspec.adapters.spanner._typing import SpannerClient as Client - if self._client is None: client_kwargs = self._connection_kwargs_for(_CLIENT_CONFIG_FIELDS) self._client = Client(**client_kwargs) @@ -392,10 +391,6 @@ def create_connection(self) -> SpannerConnection: return cast("SpannerConnection", self.get_database().snapshot(multi_use=True)) # type: ignore[no-untyped-call] def _create_pool(self) -> "AbstractSessionPool": - from sqlspec.adapters.spanner._typing import SpannerBurstyPool as BurstyPool - from sqlspec.adapters.spanner._typing import SpannerFixedSizePool as FixedSizePool - from sqlspec.adapters.spanner._typing import SpannerPingingPool as PingingPool - instance_id = self.connection_config.get("instance_id") database_id = self.connection_config.get("database_id") if not instance_id or not database_id: diff --git a/sqlspec/adapters/spanner/core.py b/sqlspec/adapters/spanner/core.py index b2c863cb0..febb40f6f 100644 --- a/sqlspec/adapters/spanner/core.py +++ b/sqlspec/adapters/spanner/core.py @@ -275,35 +275,28 @@ def create_mapped_exception(error: Any, *, logger: Any | None = None) -> SQLSpec A SQLSpec exception that wraps the original error """ del logger - # Integrity errors if isinstance(error, api_exceptions.AlreadyExists): return _create_spanner_error(error, UniqueViolationError, "resource already exists") - # Resource not found if isinstance(error, api_exceptions.NotFound): return _create_spanner_error(error, NotFoundError, "resource not found") - # SQL/argument errors if isinstance(error, api_exceptions.InvalidArgument): return _create_spanner_error(error, SQLParsingError, "invalid query or argument") - # Permission/authentication errors if isinstance(error, api_exceptions.PermissionDenied): return _create_spanner_error(error, PermissionDeniedError, "permission denied") if isinstance(error, api_exceptions.Unauthenticated): return _create_spanner_error(error, PermissionDeniedError, "authentication failed") - # Transaction errors (deadlock/abort) if isinstance(error, api_exceptions.Aborted): return _create_spanner_error(error, DeadlockError, "transaction aborted") - # Query timeout/cancellation if isinstance(error, api_exceptions.Cancelled): return _create_spanner_error(error, OperationCancelledError, "operation cancelled") if isinstance(error, api_exceptions.DeadlineExceeded): return _create_spanner_error(error, QueryTimeoutError, "deadline exceeded") - # Service/operational errors if isinstance(error, (api_exceptions.ServiceUnavailable, api_exceptions.TooManyRequests)): return _create_spanner_error(error, OperationalError, "service unavailable or rate limited") diff --git a/sqlspec/adapters/spanner/driver.py b/sqlspec/adapters/spanner/driver.py index bb195c112..0bbe76363 100644 --- a/sqlspec/adapters/spanner/driver.py +++ b/sqlspec/adapters/spanner/driver.py @@ -213,10 +213,11 @@ def dispatch_execute_script(self, cursor: "SpannerConnection", statement: "SQL") coerced_params = self._coerce_params(script_params) read_execute_kwargs = self._execute_kwargs(for_read=True) write_execute_kwargs = self._execute_kwargs() + dialect_str = str(self.dialect) if self.dialect else "spanner" for index, stmt in enumerate(statements): try: - parsed = _sqlglot.parse_one(stmt) - is_select = isinstance(parsed, _sqlglot_exp.Select) + parsed = _sqlglot.parse_one(stmt, read=dialect_str) + is_select = isinstance(parsed, _sqlglot_exp.Query) except Exception: is_select = stmt.upper().strip().startswith("SELECT") if not is_select and not is_transaction: @@ -445,7 +446,9 @@ def load_from_arrow( arrow_table = self._coerce_arrow_table(source) if overwrite: - delete_sql = f"DELETE FROM {table} WHERE TRUE" + dialect_str = str(self.dialect) if self.dialect else "spanner" + table_sql = _sqlglot_exp.to_table(table, dialect=dialect_str).sql(dialect=dialect_str, identify=True) + delete_sql = f"DELETE FROM {table_sql} WHERE TRUE" if isinstance(self.connection, SpannerTransaction): writer = cast("_SpannerWriteProtocol", self.connection) writer.execute_update(delete_sql) diff --git a/sqlspec/adapters/spanner/type_converter.py b/sqlspec/adapters/spanner/type_converter.py index 11513b4be..eda78791d 100644 --- a/sqlspec/adapters/spanner/type_converter.py +++ b/sqlspec/adapters/spanner/type_converter.py @@ -22,6 +22,8 @@ from typing import TYPE_CHECKING, Any, cast from uuid import UUID +from sqlspec.adapters.spanner._typing import SpannerJsonObject as JsonObject +from sqlspec.adapters.spanner._typing import spanner_param_types as param_types from sqlspec.core import TypedParameter from sqlspec.utils.module_loader import import_optional_attr from sqlspec.utils.type_converters import should_json_encode_sequence @@ -321,8 +323,6 @@ def _null_param_type(raw_value: Any, param_types: "SpannerParamTypesProtocol") - def _get_param_types() -> "SpannerParamTypesProtocol": global _SPANNER_PARAM_TYPES if _SPANNER_PARAM_TYPES is None: - from sqlspec.adapters.spanner._typing import spanner_param_types as param_types - _SPANNER_PARAM_TYPES = cast("SpannerParamTypesProtocol", param_types) return _SPANNER_PARAM_TYPES @@ -330,8 +330,6 @@ def _get_param_types() -> "SpannerParamTypesProtocol": def _get_json_object_type() -> "type[Any]": global _JSON_OBJECT_TYPE if _JSON_OBJECT_TYPE is None: - from sqlspec.adapters.spanner._typing import SpannerJsonObject as JsonObject - _JSON_OBJECT_TYPE = JsonObject return _JSON_OBJECT_TYPE diff --git a/sqlspec/adapters/sqlite/_typing.py b/sqlspec/adapters/sqlite/_typing.py index 25dfe5676..a8285c207 100644 --- a/sqlspec/adapters/sqlite/_typing.py +++ b/sqlspec/adapters/sqlite/_typing.py @@ -22,8 +22,7 @@ SqliteConnection: TypeAlias = _SqliteConnection SqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] SqliteRawCursor: TypeAlias = sqlite3.Cursor - -if not TYPE_CHECKING: +else: SqliteConnection = _SqliteConnection SqliteConnectionFactory = type[sqlite3.Connection] SqliteRawCursor = sqlite3.Cursor diff --git a/sqlspec/adapters/sqlite/adk/store.py b/sqlspec/adapters/sqlite/adk/store.py index 4c22f5ffd..d59d755e3 100644 --- a/sqlspec/adapters/sqlite/adk/store.py +++ b/sqlspec/adapters/sqlite/adk/store.py @@ -786,11 +786,11 @@ def create_tables(self) -> None: Skips table creation if memory store is disabled. """ - if not self.create_schema_enabled: - self.reconcile_schema() + if not self._enabled: return - if not self._enabled: + if not self.create_schema_enabled: + self.reconcile_schema() return with self._config.provide_session() as driver: @@ -836,7 +836,8 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object timestamp_julian = _datetime_to_julian(entry["timestamp"]) inserted_at_julian = _datetime_to_julian(entry["inserted_at"]) content_json_str = to_json(entry["content_json"]) - metadata_json_str = to_json(entry["metadata_json"]) if entry["metadata_json"] else None + metadata_json = entry.get("metadata_json") + metadata_json_str = to_json(metadata_json) if metadata_json is not None else None scope = entry.get("scope", "user") params_list.append(( entry["id"], @@ -845,7 +846,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["user_id"], scope, entry["event_id"], - entry["author"], + entry.get("author"), owner_id, timestamp_julian, content_json_str, @@ -864,7 +865,8 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object timestamp_julian = _datetime_to_julian(entry["timestamp"]) inserted_at_julian = _datetime_to_julian(entry["inserted_at"]) content_json_str = to_json(entry["content_json"]) - metadata_json_str = to_json(entry["metadata_json"]) if entry["metadata_json"] else None + metadata_json = entry.get("metadata_json") + metadata_json_str = to_json(metadata_json) if metadata_json is not None else None scope = entry.get("scope", "user") params_list.append(( entry["id"], @@ -873,7 +875,7 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["user_id"], scope, entry["event_id"], - entry["author"], + entry.get("author"), timestamp_julian, content_json_str, entry["content_text"], @@ -901,6 +903,9 @@ def search_entries( msg = "Memory store is disabled" raise RuntimeError(msg) + if not query or not query.strip(): + return [] + effective_limit = limit if limit is not None else self._max_results if self._use_fts: @@ -908,10 +913,19 @@ def search_entries( return self._search_entries_fts(query, app_name, user_id, effective_limit, scope_filter) except Exception as exc: logger.warning("FTS search failed; falling back to simple search: %s", exc) - return self._search_entries_simple(query, app_name, user_id, effective_limit, scope_filter) + try: + return self._search_entries_simple(query, app_name, user_id, effective_limit, scope_filter) + except sqlite3.OperationalError as exc: + if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): + return [] + raise def delete_entries_by_session(self, session_id: str) -> int: """Delete all memory entries for a specific session.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + sql = f"DELETE FROM {self._memory_table} WHERE session_id = ?" with self._config.provide_connection() as conn: @@ -924,6 +938,10 @@ def delete_entries_by_session(self, session_id: str) -> int: def delete_entries_older_than(self, days: int, app_name: "str | None" = None, scope: "str | None" = None) -> int: """Delete memory entries older than specified days.""" + if not self._enabled: + msg = "Memory store is disabled" + raise RuntimeError(msg) + cutoff_julian = _datetime_to_julian(datetime.now(timezone.utc)) - days sql = f"DELETE FROM {self._memory_table} WHERE inserted_at < ?" diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index f159e1d5b..d771eb71a 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -513,6 +513,10 @@ def _can_use_execute_many_thin_path( return False if "?" not in statement: return False + if self._resolve_dml_operation_type(statement) not in {"INSERT", "UPDATE", "DELETE"}: + return False + if "RETURNING" in statement.upper(): + return False parameter_config = config.parameter_config if parameter_config.default_parameter_style is not ParameterStyle.QMARK: diff --git a/sqlspec/builder/_base.py b/sqlspec/builder/_base.py index 278ba1396..a703fd206 100644 --- a/sqlspec/builder/_base.py +++ b/sqlspec/builder/_base.py @@ -10,8 +10,17 @@ import sqlglot from sqlglot import Dialect, exp +from sqlglot import optimizer as sqlglot_optimizer from sqlglot.dialects.dialect import DialectType from sqlglot.errors import ParseError as SQLGlotParseError +from sqlglot.optimizer import RULES +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 +from sqlglot.optimizer.qualify_columns import quote_identifiers as _quote_identifiers_rule +from sqlglot.optimizer.simplify import simplify as _simplify_rule from typing_extensions import Self from sqlspec.builder._locking import register_lock_generator @@ -837,17 +846,9 @@ 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 - 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.optimize_joins import optimize_joins as _optimize_joins_rule - from sqlglot.optimizer.pushdown_predicates import pushdown_predicates as _pushdown_predicates_rule - from sqlglot.optimizer.simplify import simplify as _simplify_rule - excluded_rules = set() if not self.optimize_joins: excluded_rules.add(_optimize_joins_rule) @@ -862,7 +863,7 @@ def _optimize_expression(self, expression: exp.Expr, *, force: bool = False) -> rules = RULES if not excluded_rules else tuple(rule for rule in RULES if rule not in excluded_rules) try: - optimized = optimize( + optimized = sqlglot_optimizer.optimize( expression, schema=cast("dict[str, object] | None", self.schema), dialect=self.dialect_name, rules=rules ) cache.put_optimized(cache_key, optimized.copy()) @@ -893,8 +894,6 @@ def _optimize_insert_with_conflict( expression.set("conflict", conflict) if optimized is expression: return expression - from sqlglot.optimizer.normalize_identifiers import normalize_identifiers as _normalize_identifiers_rule - from sqlglot.optimizer.qualify_columns import quote_identifiers as _quote_identifiers_rule dialect_name = self.dialect_name quoted_conflict = _quote_identifiers_rule( diff --git a/sqlspec/builder/_merge.py b/sqlspec/builder/_merge.py index de01d9957..1877bcf3f 100644 --- a/sqlspec/builder/_merge.py +++ b/sqlspec/builder/_merge.py @@ -776,8 +776,14 @@ def build(self, dialect: "DialectType" = None) -> "Any": dialect_name = None if dialect_name: dialect_name = dialect_name.lower() - self._move_oracle_when_conditions(dialect_name) - return super().build(dialect=dialect) + original_expression = self._expression + if dialect_name == "oracle" and original_expression is not None: + self._expression = original_expression.copy() + try: + self._move_oracle_when_conditions(dialect_name) + return super().build(dialect=dialect) + finally: + self._expression = original_expression def _move_oracle_when_conditions(self, dialect_name: str | None) -> None: """Normalize WHEN clause conditions for dialect quirks. diff --git a/sqlspec/dialects/spanner/_parsers.py b/sqlspec/dialects/spanner/_parsers.py index 3283210aa..95c3c4cbd 100644 --- a/sqlspec/dialects/spanner/_parsers.py +++ b/sqlspec/dialects/spanner/_parsers.py @@ -424,7 +424,7 @@ def attach_hints(expression: exp.Expr) -> None: comments = getattr(node, "comments", None) if not comments: continue - hint_comments = [c for c in comments if c.strip().startswith("@")] + hint_comments = [c for c in comments if c.strip().startswith("@") and "=" in c.strip()[1:]] for hc in hint_comments: hint = parse_hint_expression(hc) target_table: exp.Table | None = None diff --git a/sqlspec/extensions/adk/memory/converters.py b/sqlspec/extensions/adk/memory/converters.py index 022d0a8bb..865e4eb59 100644 --- a/sqlspec/extensions/adk/memory/converters.py +++ b/sqlspec/extensions/adk/memory/converters.py @@ -7,6 +7,9 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any +from google.adk.memory.memory_entry import MemoryEntry +from google.genai import types + from sqlspec.extensions.adk.memory._types import StoredMemory from sqlspec.utils.logging import get_logger from sqlspec.utils.uuids import uuid4 @@ -15,9 +18,7 @@ from collections.abc import Mapping, Sequence from google.adk.events.event import Event - from google.adk.memory.memory_entry import MemoryEntry from google.adk.sessions import Session - from google.genai import types __all__ = ( "event_to_memory_record", @@ -29,6 +30,7 @@ ) logger = get_logger("sqlspec.extensions.adk.memory.converters") +_UNKNOWN_SESSION_ID = "__unknown_session_id__" def extract_content_text(content: "types.Content") -> str: @@ -161,13 +163,14 @@ def memory_entry_to_record( except (ValueError, TypeError): timestamp = now + record_id = entry.id or str(uuid4()) return StoredMemory( - id=entry.id or str(uuid4()), - session_id="", + id=record_id, + session_id=_UNKNOWN_SESSION_ID, app_name=app_name, user_id=user_id, scope=scope, - event_id="", + event_id=record_id, author=entry.author or "", timestamp=timestamp, content_json=content_dict, @@ -226,9 +229,6 @@ def record_to_memory_entry(record: "StoredMemory") -> "MemoryEntry": Returns: ADK MemoryEntry object with all available fields populated. """ - from google.adk.memory.memory_entry import MemoryEntry - from google.genai import types - content = types.Content.model_validate(record["content_json"]) timestamp_str = record["timestamp"].isoformat() if record["timestamp"] else None diff --git a/sqlspec/extensions/adk/memory/service.py b/sqlspec/extensions/adk/memory/service.py index a9e8de48d..058a96bb0 100644 --- a/sqlspec/extensions/adk/memory/service.py +++ b/sqlspec/extensions/adk/memory/service.py @@ -6,6 +6,8 @@ from google.adk.memory.base_memory_service import BaseMemoryService, SearchMemoryResponse from sqlspec.extensions.adk.memory.converters import ( + _UNKNOWN_SESSION_ID, + event_to_memory_record, memory_entry_to_record, records_to_memory_entries, session_to_memory_records, @@ -107,13 +109,12 @@ async def add_events_to_memory( ``StoredMemory.metadata_json``. scope: Visibility scope ('user' or 'app'). """ - from sqlspec.extensions.adk.memory.converters import event_to_memory_record - metadata_dict = dict(custom_metadata) if custom_metadata else None + resolved_session_id = session_id or _UNKNOWN_SESSION_ID records = [] for event in events: record = event_to_memory_record( - event=event, session_id=session_id or "", app_name=app_name, user_id=user_id, scope=scope + event=event, session_id=resolved_session_id, app_name=app_name, user_id=user_id, scope=scope ) if record is not None: if metadata_dict: @@ -277,6 +278,101 @@ def add_session_to_memory(self, session: "Session", scope: str = "user") -> None "Stored %d memory entries for session %s (total events: %d)", inserted_count, session.id, len(records) ) + def add_events_to_memory( + self, + *, + app_name: str, + user_id: str, + events: "Sequence[Event]", + session_id: "str | None" = None, + custom_metadata: "Mapping[str, object] | None" = None, + scope: str = "user", + ) -> None: + """Add an explicit list of events to the memory service. + + Same Event-to-StoredMemory extraction logic as + ``add_session_to_memory``, but operates on a sequence of Events + directly (no Session wrapper needed). + + Args: + app_name: The application name for memory scope. + user_id: The user ID for memory scope. + events: The events to add to memory. + session_id: Optional session ID for memory scope/partitioning. + If None, memory entries are user-scoped only. + custom_metadata: Optional portable metadata stored in + ``StoredMemory.metadata_json``. + scope: Visibility scope ('user' or 'app'). + """ + metadata_dict = dict(custom_metadata) if custom_metadata else None + resolved_session_id = session_id or _UNKNOWN_SESSION_ID + records = [] + for event in events: + record = event_to_memory_record( + event=event, session_id=resolved_session_id, app_name=app_name, user_id=user_id, scope=scope + ) + if record is not None: + if metadata_dict: + record["metadata_json"] = metadata_dict + records.append(record) + + if not records: + logger.debug( + "No content to store for events (app=%s, user=%s, count=%d)", app_name, user_id, len(list(events)) + ) + return + + inserted_count = self._store.insert_memory_entries(records) + logger.debug( + "Stored %d memory entries from %d events (app=%s, user=%s)", inserted_count, len(records), app_name, user_id + ) + + def add_memory( + self, + *, + app_name: str, + user_id: str, + memories: "Sequence[MemoryEntry]", + custom_metadata: "Mapping[str, object] | None" = None, + scope: str = "user", + ) -> None: + """Add explicit memory items directly to the memory service. + + Each entry's ``content`` is serialized to ``content_json``, text is + extracted from ``content.parts`` for ``content_text``, and + ``custom_metadata`` merges the entry-level ``entry.custom_metadata`` + with the call-level ``custom_metadata`` parameter. + + Args: + app_name: The application name for memory scope. + user_id: The user ID for memory scope. + memories: Explicit memory items to add. + custom_metadata: Optional portable metadata for memory writes. + Merged with each entry's ``custom_metadata``. + scope: Visibility scope ('user' or 'app'). + """ + call_metadata = dict(custom_metadata) if custom_metadata else {} + records = [] + for entry in memories: + record = memory_entry_to_record( + entry=entry, app_name=app_name, user_id=user_id, extra_metadata=call_metadata, scope=scope + ) + if record is not None: + records.append(record) + + if not records: + logger.debug("No content to store for memories (app=%s, user=%s)", app_name, user_id) + return + + inserted_count = self._store.insert_memory_entries(records) + logger.debug( + "Stored %d memory entries from %d memories (app=%s, user=%s)", + inserted_count, + len(records), + app_name, + user_id, + ) + def search_memory( self, *, diff --git a/sqlspec/extensions/adk/service.py b/sqlspec/extensions/adk/service.py index 5d9fd573a..4bdebc9d4 100644 --- a/sqlspec/extensions/adk/service.py +++ b/sqlspec/extensions/adk/service.py @@ -5,6 +5,8 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, cast +from google.adk.errors import StaleSessionError +from google.adk.errors.session_not_found_error import SessionNotFoundError from google.adk.sessions.base_session_service import BaseSessionService, GetSessionConfig, ListSessionsResponse from sqlspec.extensions.adk.converters import ( @@ -276,7 +278,7 @@ async def append_event(self, session: "Session", event: "Event") -> "Event": updates the in-memory session only after persistence succeeds. Implements stale-session detection: if the session has been - modified in storage since it was last loaded, a ``ValueError`` + modified in storage since it was last loaded, a ``StaleSessionError`` is raised instead of silently overwriting. ``temp:`` keys are stripped from the persisted state snapshot so @@ -290,15 +292,13 @@ async def append_event(self, session: "Session", event: "Event") -> "Event": The appended event. Raises: - ValueError: If the session has been modified in storage since - it was loaded (stale session). + SessionNotFoundError: If the session does not exist in storage. + StaleSessionError: If the session has been modified in storage + since it was loaded (stale session). """ if event.partial: return event - # Apply temp state to in-memory session so subsequent agents in - # the same invocation can read temp values, then strip temp keys - # from the event delta before persistence. self._apply_temp_state(session, event) event = self._trim_temp_delta_state(event) @@ -309,7 +309,7 @@ async def append_event(self, session: "Session", event: "Event") -> "Event": current_record = await self._call_store("get_session", session.app_name, session.user_id, session.id) if current_record is None: msg = f"Session {session.id} not found." - raise ValueError(msg) + raise SessionNotFoundError(msg) if session._storage_update_marker is not None: # pyright: ignore[reportPrivateUsage] current_marker = compute_update_marker(current_record["update_time"]) @@ -318,13 +318,13 @@ async def append_event(self, session: "Session", event: "Event") -> "Event": "The session has been modified in storage since it was loaded. " "Please reload the session before appending more events." ) - raise ValueError(msg) + raise StaleSessionError(msg) elif current_record["update_time"].timestamp() > session.last_update_time: msg = ( "The session has been modified in storage since it was loaded. " "Please reload the session before appending more events." ) - raise ValueError(msg) + raise StaleSessionError(msg) state_delta = (event.actions.state_delta if event.actions else None) or {} app_state_delta, user_state_delta, session_state_delta = split_scoped_state(filter_temp_state(state_delta)) @@ -354,11 +354,9 @@ async def append_event(self, session: "Session", event: "Event") -> "Event": ) updated_record["state"] = merge_scoped_state(updated_record["state"], app_state, user_state) - # Use the returned record directly — saves a round-trip vs a follow-up get_session(). session.last_update_time = updated_record["update_time"].timestamp() session._storage_update_marker = compute_update_marker(updated_record["update_time"]) # pyright: ignore[reportPrivateUsage] - # Update in-memory session AFTER successful persistence self._update_session_state(session, event) session.events.append(event) diff --git a/sqlspec/migrations/base.py b/sqlspec/migrations/base.py index d836efa17..078190c93 100644 --- a/sqlspec/migrations/base.py +++ b/sqlspec/migrations/base.py @@ -16,6 +16,8 @@ from sqlspec.builder._select import Select from sqlspec.builder._update import Update from sqlspec.exceptions import MigrationError +from sqlspec.migrations.templates import MigrationTemplateSettings, build_template_settings +from sqlspec.migrations.utils import resolve_default_schema, resolve_extension_migrations_path, resolve_tracker_schema from sqlspec.migrations.version import parse_version from sqlspec.utils.logging import get_logger from sqlspec.utils.module_loader import module_to_os_path @@ -24,7 +26,6 @@ from collections.abc import Awaitable from sqlspec.config import DatabaseConfigProtocol - from sqlspec.migrations.templates import MigrationTemplateSettings from sqlspec.observability import ObservabilityRuntime __all__ = ("AppliedMigrationRecord", "BaseMigrationCommands", "BaseMigrationTracker", "LoadedMigrationMetadata") @@ -440,8 +441,6 @@ def __init__(self, config: ConfigT) -> None: Args: config: The SQLSpec configuration. """ - from sqlspec.migrations.templates import build_template_settings - self.config = config migration_config = self._get_migration_config() @@ -518,9 +517,7 @@ def _get_migration_config(self) -> "dict[str, Any]": def _resolve_default_schema(self) -> str | None: """Return the configured default migration schema.""" - from sqlspec.migrations.utils import resolve_default_schema as _resolve_default_schema - - return _resolve_default_schema(self._get_migration_config()) + return resolve_default_schema(self._get_migration_config()) def _config_supports_schemas(self) -> bool: """Return whether the bound config opts into schema-aware migrations.""" @@ -539,11 +536,9 @@ def _require_schema_support(self, default_schema: str) -> None: def _resolve_tracker_schema(self) -> str | None: """Return tracker schema only for adapters that support schema-qualified migration tables.""" - from sqlspec.migrations.utils import resolve_tracker_schema as _resolve_tracker_schema - if not self._config_supports_schemas(): return None - return _resolve_tracker_schema(self._get_migration_config()) + return resolve_tracker_schema(self._get_migration_config()) def _create_tracker(self) -> Any: """Create the configured migration tracker without breaking legacy constructors.""" @@ -584,8 +579,6 @@ def _discover_extension_migrations(self) -> "dict[str, Path]": Returns: Dictionary mapping extension names to their migration paths. """ - from sqlspec.migrations.utils import resolve_extension_migrations_path - extension_migrations = {} for ext_name, ext_options in self.extension_configs.items(): diff --git a/sqlspec/migrations/schema.py b/sqlspec/migrations/schema.py index 6d1394959..288113b47 100644 --- a/sqlspec/migrations/schema.py +++ b/sqlspec/migrations/schema.py @@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any -from sqlglot import exp, parse +from sqlglot import TokenType, exp, parse, tokenize from sqlspec.builder._ddl import AlterTable, CreateTable, _parse_ddl_identifier, _parse_ddl_table @@ -400,16 +400,28 @@ def _extract_create_table_statement(create_statement: str) -> "str | None": start = upper_statement.find("CREATE TABLE") if start < 0: return None - opening = create_statement.find("(", start) - if opening < 0: + prefix = create_statement[:start].rstrip() + if prefix.endswith("'"): + quote_pos = create_statement.rfind("'", 0, start) + try: + string_tokens = tokenize(create_statement[quote_pos:]) + except Exception: + return None + if not string_tokens or string_tokens[0].token_type != TokenType.STRING: + return None + sql = string_tokens[0].text + else: + sql = create_statement[start:] + try: + tokens = tokenize(sql) + except Exception: return None depth = 0 - for index in range(opening, len(create_statement)): - character = create_statement[index] - if character == "(": + for token in tokens: + if token.token_type == TokenType.L_PAREN: depth += 1 - elif character == ")": + elif token.token_type == TokenType.R_PAREN and depth > 0: depth -= 1 if depth == 0: - return create_statement[start : index + 1].replace("''", "'") + return sql[: token.end + 1] return None diff --git a/sqlspec/migrations/tracker.py b/sqlspec/migrations/tracker.py index 7778c6f0a..48c72a482 100644 --- a/sqlspec/migrations/tracker.py +++ b/sqlspec/migrations/tracker.py @@ -10,6 +10,7 @@ from mypy_extensions import mypyc_attr from sqlspec.migrations.base import BaseMigrationTracker +from sqlspec.migrations.schema import SchemaTarget, ensure_schema_async, ensure_schema_sync from sqlspec.observability import resolve_db_system from sqlspec.utils.logging import get_logger, log_with_context @@ -220,8 +221,6 @@ def _migrate_schema_if_needed(self, driver: "SyncDriverAdapterBase") -> None: Args: driver: The database driver to use. """ - from sqlspec.migrations.schema import SchemaTarget, ensure_schema_sync - try: target = SchemaTarget(self.version_table, self._tracking_table_ddl()) result = ensure_schema_sync(driver, [target], manage_schema=True, create_schema=False, assume_existing=True) @@ -446,8 +445,6 @@ async def _migrate_schema_if_needed(self, driver: "AsyncDriverAdapterBase") -> N Args: driver: The database driver to use. """ - from sqlspec.migrations.schema import SchemaTarget, ensure_schema_async - try: target = SchemaTarget(self.version_table, self._tracking_table_ddl()) result = await ensure_schema_async( diff --git a/tests/unit/adapters/test_aiomysql/test_adk_store.py b/tests/unit/adapters/test_aiomysql/test_adk_store.py index a84bd949d..0b970ba96 100644 --- a/tests/unit/adapters/test_aiomysql/test_adk_store.py +++ b/tests/unit/adapters/test_aiomysql/test_adk_store.py @@ -212,3 +212,103 @@ async def test_aiomysql_list_sessions_rejects_invalid_options(options: "dict[str await store.list_sessions("app", **options) assert cursor.calls == [] + + +async def test_aiomysql_adk_memory_store_insert_and_search_json_handling() -> None: + """AiomysqlADKMemoryStore handles missing author, metadata_json=None, and deserializes JSON strings.""" + from datetime import datetime, timezone + from unittest.mock import AsyncMock + + now = datetime.now(tz=timezone.utc) + cursor = MagicMock() + cursor.rowcount = 1 + cursor.execute = AsyncMock() + cursor.description = [ + ("id",), + ("session_id",), + ("app_name",), + ("user_id",), + ("scope",), + ("event_id",), + ("author",), + ("timestamp",), + ("content_json",), + ("content_text",), + ("metadata_json",), + ("inserted_at",), + ] + cursor.fetchall = AsyncMock( + return_value=[ + ( + "mem-1", + "sess-1", + "app", + "user-1", + "user", + "evt-1", + None, + now, + '{"text": "hello"}', + "hello", + '{"source": "unit"}', + now, + ) + ] + ) + cursor.close = AsyncMock() + conn = MagicMock() + conn.cursor = AsyncMock(return_value=cursor) + conn.commit = AsyncMock() + conn.__aenter__ = AsyncMock(return_value=conn) + conn.__aexit__ = AsyncMock(return_value=None) + config = _mock_config({"owner_id_column": "owner_id INT NULL"}) + config.provide_connection = lambda *_a, **_k: conn + store = AiomysqlADKMemoryStore(config) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = await store.insert_memory_entries([entry], owner_id=99) + records = await store.search_entries("hello", "app", "user-1") + + assert inserted == 1 + insert_params = cursor.execute.call_args_list[0].args[1] + assert insert_params[6] is None + assert insert_params[11] is None + assert len(records) == 1 + assert records[0]["content_json"] == {"text": "hello"} + assert records[0]["metadata_json"] == {"source": "unit"} + + +async def test_aiomysql_stream_source_closes_cursor_when_execute_raises() -> None: + """AiomysqlStreamSource.start() closes the cursor if execute() fails.""" + from unittest.mock import AsyncMock + + from sqlspec.adapters.aiomysql.core import AiomysqlStreamSource + + cursor = MagicMock() + cursor.execute = AsyncMock(side_effect=RuntimeError("execute boom")) + cursor.close = AsyncMock() + connection = MagicMock() + connection.cursor = AsyncMock(return_value=cursor) + driver = MagicMock(connection=connection) + driver._run_with_exception_handler = lambda _handler, fn: fn() + source = AiomysqlStreamSource(driver, "SELECT 1", (), 100, set()) + + with pytest.raises(RuntimeError, match="execute boom"): + await source.start() + + cursor.close.assert_awaited_once() + assert source._cursor is None diff --git a/tests/unit/adapters/test_asyncmy/test_adk_store.py b/tests/unit/adapters/test_asyncmy/test_adk_store.py index b86f87911..5b6fca77b 100644 --- a/tests/unit/adapters/test_asyncmy/test_adk_store.py +++ b/tests/unit/adapters/test_asyncmy/test_adk_store.py @@ -215,3 +215,104 @@ async def test_asyncmy_list_sessions_rejects_invalid_options(options: "dict[str, await store.list_sessions("app", **options) assert cursor.calls == [] + + +async def test_asyncmy_adk_memory_store_insert_and_search_json_handling() -> None: + """AsyncmyADKMemoryStore handles missing author, metadata_json=None, and deserializes JSON strings.""" + from datetime import datetime, timezone + from unittest.mock import AsyncMock + + now = datetime.now(tz=timezone.utc) + cursor = MagicMock() + cursor.rowcount = 1 + cursor.execute = AsyncMock() + cursor.description = [ + ("id",), + ("session_id",), + ("app_name",), + ("user_id",), + ("scope",), + ("event_id",), + ("author",), + ("timestamp",), + ("content_json",), + ("content_text",), + ("metadata_json",), + ("inserted_at",), + ] + cursor.fetchall = AsyncMock( + return_value=[ + ( + "mem-1", + "sess-1", + "app", + "user-1", + "user", + "evt-1", + None, + now, + '{"text": "hello"}', + "hello", + '{"source": "unit"}', + now, + ) + ] + ) + cursor.__aenter__ = AsyncMock(return_value=cursor) + cursor.__aexit__ = AsyncMock(return_value=None) + conn = MagicMock() + conn.cursor = MagicMock(return_value=cursor) + conn.commit = AsyncMock() + conn.__aenter__ = AsyncMock(return_value=conn) + conn.__aexit__ = AsyncMock(return_value=None) + config = _mock_config({"owner_id_column": "owner_id INT NULL"}) + config.provide_connection = lambda *_a, **_k: conn + store = AsyncmyADKMemoryStore(config) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = await store.insert_memory_entries([entry], owner_id=99) + records = await store.search_entries("hello", "app", "user-1") + + assert inserted == 1 + insert_params = cursor.execute.call_args_list[0].args[1] + assert insert_params[6] is None + assert insert_params[11] is None + assert len(records) == 1 + assert records[0]["content_json"] == {"text": "hello"} + assert records[0]["metadata_json"] == {"source": "unit"} + + +async def test_asyncmy_stream_source_closes_cursor_when_execute_raises() -> None: + """AsyncmyStreamSource.start() closes the cursor if execute() fails.""" + from unittest.mock import AsyncMock + + from sqlspec.adapters.asyncmy.core import AsyncmyStreamSource + + cursor = MagicMock() + cursor.execute = AsyncMock(side_effect=RuntimeError("execute boom")) + cursor.close = AsyncMock() + connection = MagicMock() + connection.cursor.return_value = cursor + driver = MagicMock(connection=connection) + driver._run_with_exception_handler = lambda _handler, fn: fn() + source = AsyncmyStreamSource(driver, "SELECT 1", (), 100, set()) + + with pytest.raises(RuntimeError, match="execute boom"): + await source.start() + + cursor.close.assert_awaited_once() + assert source._cursor is None diff --git a/tests/unit/adapters/test_asyncmy/test_driver.py b/tests/unit/adapters/test_asyncmy/test_driver.py index 6a78f7d8b..3ba8a226a 100644 --- a/tests/unit/adapters/test_asyncmy/test_driver.py +++ b/tests/unit/adapters/test_asyncmy/test_driver.py @@ -10,15 +10,22 @@ class _FakeConnection: - def __init__(self, in_transaction: bool) -> None: - self._in_transaction = in_transaction + def __init__(self, server_status: int) -> None: + self.server_status = server_status - def get_transaction_status(self) -> bool: - return self._in_transaction +class _StatelessConnection: + pass -@pytest.mark.parametrize("in_transaction", [True, False]) -def test_connection_in_transaction_reflects_driver_state(in_transaction: bool) -> None: - """_connection_in_transaction() must reflect the connection's real transaction status.""" - driver = AsyncmyDriver(connection=cast("Any", _FakeConnection(in_transaction))) - assert driver._connection_in_transaction() is in_transaction + +@pytest.mark.parametrize(("server_status", "expected"), [(0, False), (1, True), (2, False), (3, True)]) +def test_connection_in_transaction_reflects_driver_state(server_status: int, expected: bool) -> None: + """_connection_in_transaction() must reflect the SERVER_STATUS_IN_TRANS bit on server_status.""" + driver = AsyncmyDriver(connection=cast("Any", _FakeConnection(server_status))) + assert driver._connection_in_transaction() is expected + + +def test_connection_in_transaction_defaults_false_without_server_status() -> None: + """_connection_in_transaction() returns False when connection has no server_status attribute.""" + driver = AsyncmyDriver(connection=cast("Any", _StatelessConnection())) + assert driver._connection_in_transaction() is False diff --git a/tests/unit/adapters/test_bigquery/test_adk_store.py b/tests/unit/adapters/test_bigquery/test_adk_store.py index 14066c9b1..7c8cb2799 100644 --- a/tests/unit/adapters/test_bigquery/test_adk_store.py +++ b/tests/unit/adapters/test_bigquery/test_adk_store.py @@ -2,6 +2,7 @@ """Unit tests for BigQuery ADK store behavior.""" import inspect +from contextlib import nullcontext from datetime import datetime, timezone from typing import Any, cast, get_args, get_origin @@ -10,6 +11,7 @@ from sqlspec.adapters.bigquery import BigQueryConfig from sqlspec.adapters.bigquery.adk import BigQueryADKConfig, BigQueryADKRetentionConfig, BigQueryADKStore +from sqlspec.adapters.bigquery.litestar import BigQueryStore from sqlspec.config import ADKConfig, ExtensionConfigs from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseSyncADKStore @@ -294,3 +296,67 @@ def test_bigquery_list_sessions_rejects_invalid_options(monkeypatch: Any, option store.list_sessions("app", **options) assert calls == [] + + +def test_bigquery_litestar_store_includes_partition_filter_when_required() -> None: + """BigQueryStore appends the expires_at partition predicate when require_partition_filter is enabled.""" + executed: list[str] = [] + selected: list[str] = [] + now = datetime(2026, 5, 10, 12, 0, tzinfo=timezone.utc) + + class _FakeDriver: + def select_one(self, sql: str, **_kwargs: Any) -> dict[str, Any]: + selected.append(sql) + return {"data": b"payload", "expires_at": now} + + def execute(self, sql: str, **_kwargs: Any) -> None: + executed.append(sql) + + config = BigQueryConfig( + connection_config={"project": "proj", "dataset_id": "ds"}, + extension_config={"litestar": {"require_partition_filter": True}}, + ) + config.provide_session = lambda *_args, **_kwargs: nullcontext(_FakeDriver()) # type: ignore[method-assign] + store = BigQueryStore(config) + + store._get("s1", renew_for=60) + store._set("s1", b"val", expires_in=60) + store._delete("s1") + store._delete_all() + store._expires_in("s1") + + predicate = "expires_at IS NULL OR expires_at >= TIMESTAMP('1970-01-01 00:00:00+00')" + target_predicate = "target.expires_at IS NULL OR target.expires_at >= TIMESTAMP('1970-01-01 00:00:00+00')" + assert predicate in executed[0] + assert target_predicate in executed[1] + assert predicate in executed[2] + assert predicate in executed[3] + assert predicate in selected[1] + + +def test_bigquery_litestar_store_omits_partition_filter_by_default() -> None: + """BigQueryStore omits the synthetic partition predicate when require_partition_filter is disabled.""" + executed: list[str] = [] + selected: list[str] = [] + now = datetime(2026, 5, 10, 12, 0, tzinfo=timezone.utc) + + class _FakeDriver: + def select_one(self, sql: str, **_kwargs: Any) -> dict[str, Any]: + selected.append(sql) + return {"data": b"payload", "expires_at": now} + + def execute(self, sql: str, **_kwargs: Any) -> None: + executed.append(sql) + + config = BigQueryConfig(connection_config={"project": "proj", "dataset_id": "ds"}) + config.provide_session = lambda *_args, **_kwargs: nullcontext(_FakeDriver()) # type: ignore[method-assign] + store = BigQueryStore(config) + + store._get("s1", renew_for=60) + store._set("s1", b"val", expires_in=60) + store._delete("s1") + store._delete_all() + store._expires_in("s1") + + assert all("1970-01-01" not in sql for sql in executed) + assert all("1970-01-01" not in sql for sql in selected) diff --git a/tests/unit/adapters/test_bigquery/test_core.py b/tests/unit/adapters/test_bigquery/test_core.py index eaad631f0..ebc858abc 100644 --- a/tests/unit/adapters/test_bigquery/test_core.py +++ b/tests/unit/adapters/test_bigquery/test_core.py @@ -498,6 +498,23 @@ def test_script_execution_runs_as_single_unsplit_job() -> None: assert result.successful_statements == 1 +def test_script_execution_counts_multiple_statements() -> None: + """Verify dispatch_execute_script reports the split statement count for multi-statement scripts.""" + connection = _RecordingConnection() + script_job = SimpleNamespace( + statement_type="SCRIPT", num_dml_affected_rows=3, statistics=None, result=lambda **kwargs: None + ) + connection.job = script_job + driver = BigQueryDriver(cast(Any, connection)) + + script_sql = "INSERT INTO t VALUES (1); INSERT INTO t VALUES (2); INSERT INTO t VALUES (3);" + result = driver.dispatch_execute_script(cast(Any, connection), driver.prepare_statement(script_sql)) + + assert result.is_script_result is True + assert result.statement_count == 3 + assert result.successful_statements == 3 + + def test_script_preserves_bound_parameters() -> None: """Verify scripts retain native parameter binding without reparsing.""" connection = _RecordingConnection() diff --git a/tests/unit/adapters/test_bigquery/test_storage_write_api.py b/tests/unit/adapters/test_bigquery/test_storage_write_api.py index 576fc694d..03a01c999 100644 --- a/tests/unit/adapters/test_bigquery/test_storage_write_api.py +++ b/tests/unit/adapters/test_bigquery/test_storage_write_api.py @@ -244,8 +244,8 @@ def __init__(self, **kwargs: Any) -> None: module = SimpleNamespace(BigQueryWriteClient=_WriteClient) config = BigQueryConfig(connection_config={"project": "p", "dataset_id": "d"}) with patch.object(bigquery_config, "BigQueryStorageWriteModule", module): - first = config.provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) - second = config.provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) + first = config._provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) + second = config._provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) assert first is second assert len(built) == 1 @@ -263,7 +263,7 @@ def __init__(self, **kwargs: Any) -> None: module = SimpleNamespace(BigQueryWriteClient=_WriteClient) config = BigQueryConfig(connection_config={"project": "p", "client_options": cast("Any", options)}) with patch.object(bigquery_config, "BigQueryStorageWriteModule", module): - config.provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) + config._provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) assert built[0]["client_options"] is options @@ -280,7 +280,7 @@ def __init__(self, **_kwargs: Any) -> None: owned = BigQueryConfig(connection_config={"project": "p"}) owned._connection_instance = cast("Any", SimpleNamespace(close=lambda: closed.append("owned"))) with patch.object(bigquery_config, "BigQueryStorageWriteModule", module): - owned.provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) + owned._provide_storage_write_client(cast("Any", SimpleNamespace(_credentials=None))) owned.close_pool() assert closed == ["write", "owned"] diff --git a/tests/unit/adapters/test_db2/test_pool.py b/tests/unit/adapters/test_db2/test_pool.py index 183a65a38..573155bb7 100644 --- a/tests/unit/adapters/test_db2/test_pool.py +++ b/tests/unit/adapters/test_db2/test_pool.py @@ -144,9 +144,11 @@ def _open() -> None: worker.join() assert len({id(conn) for conn in opened}) == 2 + assert pool.size() == 2 pool.close() + assert pool.size() == 0 assert [conn.closed for conn in opened] == [True, True] diff --git a/tests/unit/adapters/test_db2/test_transactions.py b/tests/unit/adapters/test_db2/test_transactions.py index c7fec7382..1e73de987 100644 --- a/tests/unit/adapters/test_db2/test_transactions.py +++ b/tests/unit/adapters/test_db2/test_transactions.py @@ -357,3 +357,28 @@ def fail_prepare(_driver: object) -> None: assert connection.rollbacks == 1 assert connection.autocommit is True assert released == [(connection, RuntimeError)] + + +@pytest.mark.anyio +async def test_rollback_failure_resets_transaction_state_and_restores_autocommit( + fake_ibm_db: FakeModules, db2_mode: DriverMode, monkeypatch: pytest.MonkeyPatch +) -> None: + """A failing rollback must still reset _transaction_active and restore autocommit.""" + connection = _registered_connection(fake_ibm_db) + driver = db2_mode.driver(connection) + await db2_mode.call(driver.begin) + assert driver._connection_in_transaction() is True + assert connection.autocommit is False + + def failing_rollback() -> None: + raise FakeDb2OperationalError("SQL30081N A communication error has been detected.") + + monkeypatch.setattr(connection, "rollback", failing_rollback) + + with pytest.raises(SQLSpecError, match="Failed to rollback Db2 transaction"): + await db2_mode.call(driver.rollback) + + assert driver._connection_in_transaction() is False + assert connection.autocommit is True + assert not hasattr(driver, "release_open_work") + assert hasattr(driver, "_release_open_work") diff --git a/tests/unit/adapters/test_mssql_python/test_adk_store.py b/tests/unit/adapters/test_mssql_python/test_adk_store.py index 20b5cc9bf..a06106fc9 100644 --- a/tests/unit/adapters/test_mssql_python/test_adk_store.py +++ b/tests/unit/adapters/test_mssql_python/test_adk_store.py @@ -245,9 +245,9 @@ def test_mssql_python_adk_memory_store_drop_table_sql() -> None: assert store._drop_memory_table_sql() == ["DROP TABLE IF EXISTS [dbo].[adk_memory]"] -@pytest.mark.parametrize("major", [16, 17]) -def test_sync_store_defaults_to_driver_supported_json_storage(major: int) -> None: - """Server JSON availability does not imply native driver JSON support.""" +@pytest.mark.parametrize(("major", "expected_type"), [(16, "NVARCHAR(MAX)"), (17, "JSON")]) +def test_sync_store_defaults_to_driver_supported_json_storage(major: int, expected_type: str) -> None: + """When native_json is unset, DDL lazily queries server version for JSON column support.""" from sqlspec.adapters.mssql_python.data_dictionary import MssqlVersionInfo config = _mock_config() @@ -255,10 +255,53 @@ def test_sync_store_defaults_to_driver_supported_json_storage(major: int) -> Non driver.data_dictionary.get_version.return_value = MssqlVersionInfo(major=major) store = MssqlPythonADKStore(config) + assert f"state {expected_type} NOT NULL" in store._sessions_table_ddl() + config.provide_session.assert_called_once() + + +def test_sync_store_can_force_fallback_json_from_extension_config() -> None: + """Explicit native_json=False uses NVARCHAR(MAX) without opening a session.""" + config = _mock_config({"native_json": False}) + store = MssqlPythonADKStore(config) + assert "state NVARCHAR(MAX) NOT NULL" in store._sessions_table_ddl() config.provide_session.assert_not_called() +def test_mssql_python_adk_memory_store_insert_handles_none_metadata_and_missing_author() -> None: + """insert_memory_entries binds None for metadata_json=None and missing author.""" + from datetime import datetime, timezone + + config = _mock_config() + conn = config.provide_connection.return_value.__enter__.return_value + cursor = conn.cursor.return_value + cursor.rowcount = 1 + store = MssqlPythonADKMemoryStore(config) + now = datetime.now(tz=timezone.utc) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = store.insert_memory_entries([entry]) + + assert inserted == 1 + params = cursor.execute.call_args.args[1] + assert params[6] is None + assert params[10] is None + + def test_disabled_memory_store_rejects_operations_without_connecting() -> None: config = _mock_config({"enable_memory": False}) store = MssqlPythonADKMemoryStore(config) diff --git a/tests/unit/adapters/test_mssql_python/test_config.py b/tests/unit/adapters/test_mssql_python/test_config.py index 4d35a8473..a96c1d3db 100644 --- a/tests/unit/adapters/test_mssql_python/test_config.py +++ b/tests/unit/adapters/test_mssql_python/test_config.py @@ -359,3 +359,16 @@ def test_pool_does_not_reconfigure_when_params_match(monkeypatch: pytest.MonkeyP MssqlPythonConnectionPool(connection_string="Server=localhost;", max_size=10, idle_timeout=60, enabled=True) assert not any("Pooling configuration was already set" in str(w.message) for w in recorded) assert calls == [] + + +def test_close_pool_clears_connection_instance(monkeypatch: pytest.MonkeyPatch) -> None: + """MssqlPythonConfig._close_pool closes the pool and clears connection_instance.""" + monkeypatch.setattr(_mssql_pool, "_POOLING_PARAMS", None) + monkeypatch.setattr("sqlspec.adapters.mssql_python.pool.MSSQL_PYTHON_MODULE.pooling", lambda **_: None) + config = MssqlPythonConfig(connection_config={"server": "localhost"}) + pool = config.create_pool() + assert config.connection_instance is pool + + config.close_pool() + + assert config.connection_instance is None diff --git a/tests/unit/adapters/test_mssql_python/test_core.py b/tests/unit/adapters/test_mssql_python/test_core.py index d6a1a1f86..92fa5f558 100644 --- a/tests/unit/adapters/test_mssql_python/test_core.py +++ b/tests/unit/adapters/test_mssql_python/test_core.py @@ -258,17 +258,23 @@ def test_parse_odbc_connection_string_edge_cases() -> None: def test_extract_error_number_from_attribute() -> None: - """extract_error_number retrieves native integer attribute 'number'.""" + """extract_error_number retrieves native integer attribute 'number' and skips 0.""" class CustomError(Exception): number = 2627 + class ZeroError(Exception): + number = 0 + assert extract_error_number(CustomError("duplicate key")) == 2627 + assert extract_error_number(ZeroError("Msg 2627, Level 14, State 1")) == 2627 + assert extract_error_number(ZeroError("Plain error")) is None def test_extract_error_number_from_args_tuple() -> None: - """extract_error_number extracts integer from exception args.""" + """extract_error_number extracts integer from exception args and skips 0.""" assert extract_error_number(Exception(1205, "Deadlock found")) == 1205 + assert extract_error_number(Exception(0, "Msg 1205, Level 13, State 1")) == 1205 def test_extract_error_number_from_string_regex() -> None: diff --git a/tests/unit/adapters/test_mssql_python/test_transaction_state.py b/tests/unit/adapters/test_mssql_python/test_transaction_state.py index c6ddf6cc3..9b65db305 100644 --- a/tests/unit/adapters/test_mssql_python/test_transaction_state.py +++ b/tests/unit/adapters/test_mssql_python/test_transaction_state.py @@ -131,22 +131,36 @@ def test_mssql_python_begin_failure_keeps_transaction_inactive() -> None: assert driver._connection_in_transaction() is False # pyright: ignore[reportPrivateUsage] -@pytest.mark.parametrize("method_name", ["commit", "rollback"]) -def test_mssql_python_completion_failure_preserves_active_state(method_name: str) -> None: - """A failed DBAPI completion should leave the transaction active and un-restored.""" +def test_mssql_python_commit_failure_preserves_active_state() -> None: + """A failed DBAPI commit should leave the transaction active and un-restored.""" connection = FakeConnection(autocommit=True) driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) driver.begin() - setattr(connection, f"fail_{method_name}", True) + connection.fail_commit = True - with pytest.raises(SQLSpecError, match=f"Failed to {method_name} transaction"): - getattr(driver, method_name)() + with pytest.raises(SQLSpecError, match="Failed to commit transaction"): + driver.commit() assert connection.autocommit is False assert connection.autocommit_values == [False] assert driver._connection_in_transaction() is True # pyright: ignore[reportPrivateUsage] +def test_mssql_python_rollback_failure_resets_active_state_and_restores_autocommit() -> None: + """A failed DBAPI rollback must still reset transaction state and restore autocommit.""" + connection = FakeConnection(autocommit=True) + driver = MssqlPythonDriver(cast("MssqlPythonConnection", connection)) + driver.begin() + connection.fail_rollback = True + + with pytest.raises(SQLSpecError, match="Failed to rollback transaction"): + driver.rollback() + + assert connection.autocommit is True + assert connection.autocommit_values == [False, True] + assert driver._connection_in_transaction() is False # pyright: ignore[reportPrivateUsage] + + @pytest.mark.parametrize("method_name", ["commit", "rollback"]) def test_mssql_python_restore_failure_reports_inactive_transaction(method_name: str) -> None: """A restoration failure should surface after the DBAPI transaction has completed.""" diff --git a/tests/unit/adapters/test_mysqlconnector/test_adk_store.py b/tests/unit/adapters/test_mysqlconnector/test_adk_store.py index 259ee706d..82dd3d582 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_adk_store.py +++ b/tests/unit/adapters/test_mysqlconnector/test_adk_store.py @@ -323,3 +323,153 @@ def test_mysqlconnector_sync_list_sessions_rejects_invalid_options(options: "dic store.list_sessions("app", **options) assert cursor.calls == [] + + +def test_mysqlconnector_sync_adk_memory_store_insert_and_search_json_handling() -> None: + """MysqlConnectorSyncADKMemoryStore handles missing author, metadata_json=None, and deserializes JSON strings.""" + from datetime import datetime, timezone + + now = datetime.now(tz=timezone.utc) + cursor = MagicMock() + cursor.rowcount = 1 + cursor.description = [ + ("id",), + ("session_id",), + ("app_name",), + ("user_id",), + ("scope",), + ("event_id",), + ("author",), + ("timestamp",), + ("content_json",), + ("content_text",), + ("metadata_json",), + ("inserted_at",), + ] + cursor.fetchall.return_value = [ + ( + "mem-1", + "sess-1", + "app", + "user-1", + "user", + "evt-1", + None, + now, + '{"text": "hello"}', + "hello", + '{"source": "unit"}', + now, + ) + ] + conn = MagicMock() + conn.cursor.return_value = cursor + conn.__enter__.return_value = conn + conn.__exit__.return_value = None + config = _mock_config({"owner_id_column": "owner_id INT NULL"}) + config.provide_connection = lambda *_a, **_k: conn + store = MysqlConnectorSyncADKMemoryStore(config) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = store.insert_memory_entries([entry], owner_id=99) + records = store.search_entries("hello", "app", "user-1") + + assert inserted == 1 + insert_params = cursor.execute.call_args_list[0].args[1] + assert insert_params[6] is None + assert insert_params[11] is None + assert len(records) == 1 + assert records[0]["content_json"] == {"text": "hello"} + assert records[0]["metadata_json"] == {"source": "unit"} + + +async def test_mysqlconnector_async_adk_memory_store_insert_and_search_json_handling() -> None: + """MysqlConnectorAsyncADKMemoryStore handles missing author, metadata_json=None, and deserializes JSON strings.""" + from datetime import datetime, timezone + from unittest.mock import AsyncMock + + now = datetime.now(tz=timezone.utc) + cursor = MagicMock() + cursor.rowcount = 1 + cursor.execute = AsyncMock() + cursor.description = [ + ("id",), + ("session_id",), + ("app_name",), + ("user_id",), + ("scope",), + ("event_id",), + ("author",), + ("timestamp",), + ("content_json",), + ("content_text",), + ("metadata_json",), + ("inserted_at",), + ] + cursor.fetchall = AsyncMock( + return_value=[ + ( + "mem-1", + "sess-1", + "app", + "user-1", + "user", + "evt-1", + None, + now, + '{"text": "hello"}', + "hello", + '{"source": "unit"}', + now, + ) + ] + ) + cursor.close = AsyncMock() + conn = MagicMock() + conn.cursor = AsyncMock(return_value=cursor) + conn.commit = AsyncMock() + conn.__aenter__ = AsyncMock(return_value=conn) + conn.__aexit__ = AsyncMock(return_value=None) + config = _mock_config({"owner_id_column": "owner_id INT NULL"}) + config.provide_connection = lambda *_a, **_k: conn + store = MysqlConnectorAsyncADKMemoryStore(config) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = await store.insert_memory_entries([entry], owner_id=99) + records = await store.search_entries("hello", "app", "user-1") + + assert inserted == 1 + insert_params = cursor.execute.call_args_list[0].args[1] + assert insert_params[6] is None + assert insert_params[11] is None + assert len(records) == 1 + assert records[0]["content_json"] == {"text": "hello"} + assert records[0]["metadata_json"] == {"source": "unit"} diff --git a/tests/unit/adapters/test_mysqlconnector/test_config.py b/tests/unit/adapters/test_mysqlconnector/test_config.py index 78f211844..a693c566f 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_config.py +++ b/tests/unit/adapters/test_mysqlconnector/test_config.py @@ -453,3 +453,21 @@ async def test_async_pool_unavailable_preserves_standalone_connections(monkeypat assert await MysqlConnectorAsyncConfig()._acquire_async_connection() is connection with pytest.raises(ImproperConfigurationError, match=r"9\.4"): await MysqlConnectorAsyncConfig(connection_config={"pool_size": 2})._acquire_async_connection() + + +def test_sync_create_pool_preserves_zero_pool_size(monkeypatch: pytest.MonkeyPatch) -> None: + """MysqlConnectorSyncConfig._create_pool must preserve explicit pool_size=0.""" + from sqlspec.adapters.mysqlconnector import config as cfg_module + + pool_kwargs: list[dict[str, Any]] = [] + + class _FakePool: + def __init__(self, **kwargs: Any) -> None: + pool_kwargs.append(kwargs) + + monkeypatch.setattr(cfg_module, "MysqlConnectorConnectionPool", _FakePool) + config = MysqlConnectorSyncConfig(connection_config={"pool_size": 0}) + config._create_pool() + + assert len(pool_kwargs) == 1 + assert pool_kwargs[0]["pool_size"] == 0 diff --git a/tests/unit/adapters/test_mysqlconnector/test_core.py b/tests/unit/adapters/test_mysqlconnector/test_core.py index 1f0a53cd1..bfbcb0c1f 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_core.py +++ b/tests/unit/adapters/test_mysqlconnector/test_core.py @@ -1,6 +1,14 @@ """mysql-connector compiled core helpers.""" -from sqlspec.adapters.mysqlconnector.core import collect_stream_rows +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from sqlspec.adapters.mysqlconnector.core import ( + MysqlConnectorAsyncStreamSource, + MysqlConnectorSyncStreamSource, + collect_stream_rows, +) from sqlspec.utils.serializers import from_json @@ -18,3 +26,37 @@ def test_collect_stream_rows_still_zips_tuple_rows() -> None: collected = collect_stream_rows([(1, '{"name": "alpha"}')], (["id", "payload"], [1]), from_json) assert collected == [{"id": 1, "payload": {"name": "alpha"}}] + + +def test_sync_stream_source_closes_cursor_when_execute_raises() -> None: + """MysqlConnectorSyncStreamSource.start() closes the cursor if execute() fails.""" + cursor = MagicMock() + cursor.execute.side_effect = RuntimeError("execute boom") + connection = MagicMock() + connection.cursor.return_value = cursor + driver = MagicMock(connection=connection, driver_features={}) + source = MysqlConnectorSyncStreamSource(driver, "SELECT 1", (), 100, set()) + + with pytest.raises(RuntimeError, match="execute boom"): + source.start() + + cursor.close.assert_called_once() + assert source._cursor is None + + +async def test_async_stream_source_closes_cursor_when_execute_raises() -> None: + """MysqlConnectorAsyncStreamSource.start() closes the cursor if execute() fails.""" + cursor = MagicMock() + cursor.execute = AsyncMock(side_effect=RuntimeError("execute boom")) + cursor.close = AsyncMock() + connection = MagicMock(_cnx=None, unread_result=False) + connection.cursor = AsyncMock(return_value=cursor) + driver = MagicMock(connection=connection, driver_features={}) + driver._run_with_exception_handler = lambda _handler, fn: fn() + source = MysqlConnectorAsyncStreamSource(driver, "SELECT 1", (), 100, set()) + + with pytest.raises(RuntimeError, match="execute boom"): + await source.start() + + cursor.close.assert_awaited_once() + assert source._cursor is None diff --git a/tests/unit/adapters/test_oracledb/test_litestar_store.py b/tests/unit/adapters/test_oracledb/test_litestar_store.py index e171d38e3..9f97d52d5 100644 --- a/tests/unit/adapters/test_oracledb/test_litestar_store.py +++ b/tests/unit/adapters/test_oracledb/test_litestar_store.py @@ -3,6 +3,8 @@ from datetime import datetime, timedelta, timezone from typing import Any, cast +from typing_extensions import Self + from sqlspec.adapters.oracledb.core import DB_TYPE_BLOB from sqlspec.adapters.oracledb.data_dictionary import OracleVersionCache from sqlspec.adapters.oracledb.litestar import OracleSyncStore @@ -15,6 +17,16 @@ def __init__(self, rows: "list[tuple[Any, ...]] | None" = None) -> None: self.rows = list(rows or []) self.executed: list[tuple[str, dict[str, Any] | None]] = [] self.rowcount = 0 + self.closed = False + + def __enter__(self) -> "Self": + return self + + def __exit__(self, *_: object) -> None: + self.close() + + def close(self) -> None: + self.closed = True def execute(self, sql: str, parameters: "dict[str, Any] | None" = None) -> None: self.executed.append((sql, parameters)) diff --git a/tests/unit/adapters/test_oracledb/test_lob_coercion.py b/tests/unit/adapters/test_oracledb/test_lob_coercion.py index 22266d5d3..09aeda895 100644 --- a/tests/unit/adapters/test_oracledb/test_lob_coercion.py +++ b/tests/unit/adapters/test_oracledb/test_lob_coercion.py @@ -82,12 +82,13 @@ def test_coerce_json_parameters_sync_pre_native_versions_create_utf8_blob_locato def test_coerce_json_parameters_sync_native_versions_keep_python_value_for_db_type_json( sync_connection: MagicMock, payload: object, wrapper: "Callable[[object], object]" ) -> None: - """Oracle 21c+ values stay as Python JSON for the DB_TYPE_JSON input handler.""" + """Oracle 21c+ values stay as Python JSON or OracleJson for the DB_TYPE_JSON input handler.""" sync_connection._sqlspec_oracle_major = 21 + wrapped = wrapper(payload) result = coerce_large_parameters_sync( sync_connection, - {"payload": wrapper(payload)}, + {"payload": wrapped}, clob_type=CLOB_TYPE, blob_type=BLOB_TYPE, varchar2_byte_limit=VARCHAR2_LIMIT, @@ -95,7 +96,7 @@ def test_coerce_json_parameters_sync_native_versions_keep_python_value_for_db_ty ) sync_connection.createlob.assert_not_called() - assert result["payload"] is payload + assert result["payload"] is wrapped def test_coerce_json_parameters_sync_uses_connection_version_when_cached_major_is_missing( @@ -388,7 +389,7 @@ def test_coerce_large_parameters_sync_string_exactly_at_threshold_no_coercion(sy def test_coerce_large_parameters_sync_string_over_threshold_becomes_clob(sync_connection: MagicMock) -> None: params = {"content": "a" * 4001} - coerce_large_parameters_sync( + result = coerce_large_parameters_sync( sync_connection, params, clob_type=CLOB_TYPE, @@ -397,7 +398,9 @@ def test_coerce_large_parameters_sync_string_over_threshold_becomes_clob(sync_co raw_byte_limit=RAW_LIMIT, ) sync_connection.createlob.assert_called_once_with(CLOB_TYPE, "a" * 4001) - assert params["content"] is sync_connection.createlob.return_value + assert result["content"] is sync_connection.createlob.return_value + assert result is not params + assert params["content"] == "a" * 4001 def test_coerce_large_parameters_sync_multibyte_string_under_charcount_but_over_bytecount( @@ -466,7 +469,7 @@ def test_coerce_large_parameters_sync_mixed_parameters(sync_connection: MagicMoc "big_bytes": b"\xff" * 3000, "number": 42, } - coerce_large_parameters_sync( + result = coerce_large_parameters_sync( sync_connection, params, clob_type=CLOB_TYPE, @@ -474,10 +477,11 @@ def test_coerce_large_parameters_sync_mixed_parameters(sync_connection: MagicMoc varchar2_byte_limit=VARCHAR2_LIMIT, raw_byte_limit=RAW_LIMIT, ) - assert params["small_str"] == "hello" - assert params["big_str"] is sync_connection.createlob.return_value - assert params["small_bytes"] == b"\x00" * 100 - assert params["number"] == 42 + assert result["small_str"] == "hello" + assert result["big_str"] is sync_connection.createlob.return_value + assert result["small_bytes"] == b"\x00" * 100 + assert result["number"] == 42 + assert params["big_str"] == "x" * 5000 assert sync_connection.createlob.call_count == 2 @@ -488,7 +492,7 @@ def test_coerce_large_parameters_sync_oracle_clob_wrapper_short_value_routed_to_ from sqlspec.adapters.oracledb import OracleClob params = {"v": OracleClob("short text")} - coerce_large_parameters_sync( + result = coerce_large_parameters_sync( sync_connection, params, clob_type=CLOB_TYPE, @@ -497,7 +501,8 @@ def test_coerce_large_parameters_sync_oracle_clob_wrapper_short_value_routed_to_ raw_byte_limit=RAW_LIMIT, ) sync_connection.createlob.assert_called_once_with(CLOB_TYPE, "short text") - assert params["v"] is sync_connection.createlob.return_value + assert result["v"] is sync_connection.createlob.return_value + assert isinstance(params["v"], OracleClob) def test_coerce_large_parameters_sync_oracle_clob_wrapper_bytes_decoded_to_str(sync_connection: MagicMock) -> None: @@ -551,10 +556,12 @@ def test_coerce_large_parameters_sync_oracle_blob_wrapper_str_encoded_to_bytes(s def test_coerce_large_parameters_sync_oracle_json_wrapper_unwrapped_to_value(sync_connection: MagicMock) -> None: - """OracleJson unwraps so the C1 input handler can claim the value.""" + """OracleJson stays wrapped on 21c+ so the C1 input handler can claim and unwrap the value.""" from sqlspec.adapters.oracledb import OracleJson + from sqlspec.adapters.oracledb._json_handlers import is_json_payload, json_converter_in_native - params = {"v": OracleJson({"a": 1})} + wrapped = OracleJson({"a": 1}) + params = {"v": wrapped} result = coerce_large_parameters_sync( sync_connection, params, @@ -563,7 +570,9 @@ def test_coerce_large_parameters_sync_oracle_json_wrapper_unwrapped_to_value(syn varchar2_byte_limit=VARCHAR2_LIMIT, raw_byte_limit=RAW_LIMIT, ) - assert result["v"] == {"a": 1} + assert result["v"] is wrapped + assert is_json_payload(result["v"]) is True + assert json_converter_in_native(result["v"]) == {"a": 1} sync_connection.createlob.assert_not_called() @@ -628,10 +637,11 @@ def test_coerce_large_parameters_sync_oracle_blob_wrapper_unwrapped_in_positiona def test_coerce_large_parameters_sync_oracle_json_wrapper_unwrapped_in_positional_tuple( sync_connection: MagicMock, ) -> None: - """OracleJson inside a positional tuple is unwrapped to its inner value.""" + """OracleJson inside a positional tuple is preserved on 21c+ for the C1 input handler.""" from sqlspec.adapters.oracledb import OracleJson - params = (1, OracleJson({"a": 1})) + wrapped = OracleJson({"a": 1}) + params = (1, wrapped) result = coerce_large_parameters_sync( sync_connection, params, @@ -641,7 +651,7 @@ def test_coerce_large_parameters_sync_oracle_json_wrapper_unwrapped_in_positiona raw_byte_limit=RAW_LIMIT, ) sync_connection.createlob.assert_not_called() - assert result[1] == {"a": 1} + assert result[1] is wrapped def test_coerce_large_parameters_sync_positional_tuple_str_over_threshold_becomes_clob( @@ -709,7 +719,7 @@ async def test_coerce_large_parameters_async_none_parameters_passthrough(async_c @pytest.mark.anyio async def test_coerce_large_parameters_async_string_over_threshold_becomes_clob(async_connection: AsyncMock) -> None: params = {"content": "a" * 4001} - await coerce_large_parameters_async( + result = await coerce_large_parameters_async( async_connection, params, clob_type=CLOB_TYPE, @@ -718,6 +728,9 @@ async def test_coerce_large_parameters_async_string_over_threshold_becomes_clob( raw_byte_limit=RAW_LIMIT, ) async_connection.createlob.assert_called_once_with(CLOB_TYPE, "a" * 4001) + assert result["content"] is async_connection.createlob.return_value + assert result is not params + assert params["content"] == "a" * 4001 @pytest.mark.anyio @@ -829,10 +842,11 @@ async def test_coerce_large_parameters_async_oracle_blob_wrapper_str_encoded_to_ async def test_coerce_large_parameters_async_oracle_json_wrapper_unwrapped_to_value( async_connection: AsyncMock, ) -> None: - """OracleJson unwraps so the C1 input handler can claim the value.""" + """OracleJson stays wrapped on 21c+ so the C1 input handler can claim and unwrap the value.""" from sqlspec.adapters.oracledb import OracleJson - params = {"v": OracleJson({"a": 1})} + wrapped = OracleJson({"a": 1}) + params = {"v": wrapped} result = await coerce_large_parameters_async( async_connection, params, @@ -841,7 +855,7 @@ async def test_coerce_large_parameters_async_oracle_json_wrapper_unwrapped_to_va varchar2_byte_limit=VARCHAR2_LIMIT, raw_byte_limit=RAW_LIMIT, ) - assert result["v"] == {"a": 1} + assert result["v"] is wrapped async_connection.createlob.assert_not_called() @@ -906,19 +920,20 @@ async def test_coerce_large_parameters_async_oracle_blob_wrapper_unwrapped_in_po async def test_coerce_large_parameters_async_oracle_json_wrapper_unwrapped_in_positional_tuple( async_connection: AsyncMock, ) -> None: - """OracleJson inside a positional tuple is unwrapped to its inner value.""" + """OracleJson inside a positional tuple is preserved on 21c+ for the C1 input handler.""" from sqlspec.adapters.oracledb import OracleJson + wrapped = OracleJson({"a": 1}) result = await coerce_large_parameters_async( async_connection, - (1, OracleJson({"a": 1})), + (1, wrapped), clob_type=CLOB_TYPE, blob_type=BLOB_TYPE, varchar2_byte_limit=VARCHAR2_LIMIT, raw_byte_limit=RAW_LIMIT, ) async_connection.createlob.assert_not_called() - assert result[1] == {"a": 1} + assert result[1] is wrapped @pytest.mark.anyio @@ -1080,3 +1095,35 @@ async def test_async_stream_keeps_locators_when_fetch_lobs_is_requested() -> Non assert chunk == [{"id": 1, "body": locator}] assert locator.read_count == 0 + + +def test_oracle_sync_stream_source_closes_cursor_when_execute_raises() -> None: + """OracleSyncStreamSource.start() closes the cursor if execute() fails.""" + cursor = MagicMock() + cursor.execute.side_effect = RuntimeError("execute boom") + connection = MagicMock() + connection.cursor.return_value = cursor + driver = OracleSyncDriver(cast("OracleSyncConnection", connection)) + source = OracleSyncStreamSource(driver, "SELECT 1 FROM dual", None, 100) + + with pytest.raises(RuntimeError, match="execute boom"): + source.start() + + cursor.close.assert_called_once() + assert source._cursor is None + + +async def test_oracle_async_stream_source_closes_cursor_when_execute_raises() -> None: + """OracleAsyncStreamSource.start() closes the cursor if execute() fails.""" + cursor = MagicMock() + cursor.execute = AsyncMock(side_effect=RuntimeError("execute boom")) + connection = MagicMock() + connection.cursor.return_value = cursor + driver = OracleAsyncDriver(cast("OracleAsyncConnection", connection)) + source = OracleAsyncStreamSource(driver, "SELECT 1 FROM dual", None, 100) + + with pytest.raises(RuntimeError, match="execute boom"): + await source.start() + + cursor.close.assert_called_once() + assert source._cursor is None diff --git a/tests/unit/adapters/test_oracledb/test_oracle_adk_store.py b/tests/unit/adapters/test_oracledb/test_oracle_adk_store.py index 850db9820..e84e3618a 100644 --- a/tests/unit/adapters/test_oracledb/test_oracle_adk_store.py +++ b/tests/unit/adapters/test_oracledb/test_oracle_adk_store.py @@ -349,6 +349,16 @@ class _RecordingCursor: def __init__(self) -> None: self.calls: list[tuple[str, dict[str, Any]]] = [] + self.closed = False + + def __enter__(self) -> "Self": + return self + + def __exit__(self, *_: Any) -> None: + self.close() + + def close(self) -> None: + self.closed = True def execute(self, sql: str, params: "dict[str, Any] | None" = None) -> None: self.calls.append((sql, dict(params or {}))) @@ -477,6 +487,30 @@ def test_oracle_sync_list_sessions_zero_limit_never_queries() -> None: assert cursor.calls == [] +@pytest.mark.parametrize("storage_type", list(JSONStorageType)) +def test_oracle_adk_memory_table_ddl_has_no_duplicate_columns(storage_type: JSONStorageType) -> None: + """Memory table DDL must not declare duplicate app_name or user_id columns.""" + config = _mock_config({}) + for store in (OracleAsyncADKMemoryStore(config), OracleSyncADKMemoryStore(config)): + sql = store._memory_table_ddl_for_type(storage_type) + assert sql.count("app_name VARCHAR2(128) NOT NULL") == 1 + assert sql.count("user_id VARCHAR2(128) NOT NULL") == 1 + + +def test_oracle_sync_list_sessions_closes_cursor() -> None: + """Sync ADK store closes cursor via context manager.""" + store, cursor = _sync_session_store() + store.list_sessions("app") + assert cursor.closed is True + + +async def test_oracle_async_list_sessions_closes_cursor() -> None: + """Async ADK store closes cursor via context manager.""" + store, cursor = _async_session_store() + await store.list_sessions("app") + assert cursor.closed is True + + @pytest.mark.parametrize( "options", [ diff --git a/tests/unit/adapters/test_psqlpy/test_core.py b/tests/unit/adapters/test_psqlpy/test_core.py index 6c456cb17..ce9ffcf9f 100644 --- a/tests/unit/adapters/test_psqlpy/test_core.py +++ b/tests/unit/adapters/test_psqlpy/test_core.py @@ -279,11 +279,8 @@ def test_dml_count_query_wraps_supported_statements(sql: str) -> None: def test_dml_count_query_preserves_placeholders_quotes_and_existing_with() -> None: - """The rewrite should preserve compiled placeholders and a DML-owned WITH clause.""" - sql = ( - 'WITH source AS (SELECT $2 AS "id") ' - 'UPDATE "events" SET "payload" = $1 FROM source WHERE "events"."id" = source."id"' - ) + """The rewrite should preserve compiled placeholders and a nested subquery WITH clause.""" + sql = 'INSERT INTO "events" ("id", "payload") SELECT "id", $1 FROM (WITH source AS (SELECT $2 AS "id") SELECT "id" FROM source) AS sub' rewritten = psqlpy_core._dml_count_query(sql) # pyright: ignore[reportPrivateUsage] @@ -297,9 +294,8 @@ def test_dml_count_query_preserves_placeholders_quotes_and_existing_with() -> No def test_dml_count_query_uses_collision_free_cte_alias() -> None: """A user CTE using the private base name should force a deterministic suffix.""" sql = ( - "WITH _sqlspec_affected AS (SELECT $2 AS id) " - "UPDATE events SET payload = $1 FROM _sqlspec_affected " - "WHERE events.id = _sqlspec_affected.id" + "INSERT INTO events (id, payload) " + "SELECT id, $1 FROM (WITH _sqlspec_affected AS (SELECT $2 AS id) SELECT id FROM _sqlspec_affected) AS sub" ) rewritten = psqlpy_core._dml_count_query(sql) # pyright: ignore[reportPrivateUsage] @@ -316,6 +312,7 @@ def test_dml_count_query_uses_collision_free_cte_alias() -> None: "MERGE INTO events USING source ON events.id = source.id WHEN MATCHED THEN DELETE", "CREATE TABLE events (id INT)", "UPDATE events SET payload = $1 WHERE id = $2 RETURNING id", + "WITH source AS (SELECT $2 AS id) UPDATE events SET payload = $1 FROM source WHERE events.id = source.id", ], ) def test_dml_count_query_bypasses_unsupported_or_returning_statements(sql: str) -> None: @@ -324,6 +321,5 @@ def test_dml_count_query_bypasses_unsupported_or_returning_statements(sql: str) def test_dml_count_query_surfaces_parse_errors() -> None: - """Invalid compiled SQL should raise rather than report a false row count.""" - with pytest.raises(SQLSpecError, match="Unable to build psqlpy DML row count query"): - psqlpy_core._dml_count_query("UPDATE events SET payload =") # pyright: ignore[reportPrivateUsage] + """Invalid compiled SQL should return None rather than raising during count query rewrite.""" + assert psqlpy_core._dml_count_query("UPDATE events SET payload =") is None # pyright: ignore[reportPrivateUsage] diff --git a/tests/unit/adapters/test_psycopg/test_adk_store.py b/tests/unit/adapters/test_psycopg/test_adk_store.py index e526c4454..693518c1a 100644 --- a/tests/unit/adapters/test_psycopg/test_adk_store.py +++ b/tests/unit/adapters/test_psycopg/test_adk_store.py @@ -86,6 +86,7 @@ async def fetchall(self) -> "list[dict[str, Any]]": class _DummyAsyncConnection: def __init__(self, cursor: _DummyAsyncCursor) -> None: self._cursor = cursor + self.commit_called = False async def __aenter__(self) -> Self: return self @@ -96,6 +97,9 @@ async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: def cursor(self, **kwargs: Any) -> _DummyAsyncCursor: return self._cursor + async def commit(self) -> None: + self.commit_called = True + class _DummyConfig: def __init__(self, connection: _DummyConnection | _DummyAsyncConnection) -> None: @@ -407,6 +411,9 @@ async def execute(self, query: Any, params: Any = None) -> None: # type: ignore async def fetchall(self) -> "list[dict[str, Any]]": # type: ignore[override] return self._rows + async def fetchone(self) -> "dict[str, Any] | None": # type: ignore[override] + return self._rows[0] if self._rows else None + class _AsyncDummyConnection(_DummyConnection): async def __aenter__(self) -> Self: @@ -415,6 +422,9 @@ async def __aenter__(self) -> Self: async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: return None + async def commit(self) -> None: # type: ignore[override] + self.commit_called = True + def _rendered(query: Any) -> str: return " ".join(query.as_string(None).split()) @@ -511,3 +521,43 @@ def test_psycopg_sync_list_sessions_rejects_invalid_options(options: "dict[str, store.list_sessions("app", **options) assert cursor.execute_calls == [] + + +async def test_psycopg_async_adk_store_commits_mutating_operations() -> None: + """Async ADK session and memory mutating operations must commit their transactions.""" + cursor = _AsyncDummyCursor() + connection = _AsyncDummyConnection(cursor) + config = _mock_config() + config.provide_connection = lambda *_a, **_k: connection + store = PsycopgAsyncADKStore(config) + + await store.update_session_state("app", "u1", "s1", {"k": "v"}) + assert connection.commit_called + + connection.commit_called = False + await store.delete_session("app", "u1", "s1") + assert connection.commit_called + + connection.commit_called = False + await store.upsert_app_state("app", {"k": "v"}) + assert connection.commit_called + + connection.commit_called = False + await store.upsert_user_state("app", "u1", {"k": "v"}) + assert connection.commit_called + + connection.commit_called = False + await store.set_metadata("k", "v") + assert connection.commit_called + + +def test_psycopg_sync_adk_store_commits_mutating_operations() -> None: + """Sync ADK session and memory mutating operations must commit their transactions.""" + store, _, connection = _build_store() + + store.update_session_state("app", "u1", "s1", {"k": "v"}) + assert connection.commit_called + + connection.commit_called = False + store.delete_session("app", "u1", "s1") + assert connection.commit_called diff --git a/tests/unit/adapters/test_pymssql/test_adk_store.py b/tests/unit/adapters/test_pymssql/test_adk_store.py index 4f92cab91..cafdc9e87 100644 --- a/tests/unit/adapters/test_pymssql/test_adk_store.py +++ b/tests/unit/adapters/test_pymssql/test_adk_store.py @@ -97,3 +97,40 @@ def test_pymssql_list_sessions_rejects_invalid_options( store.list_sessions("app", **options) assert calls == [] + + +def test_pymssql_adk_memory_store_insert_handles_none_metadata_and_missing_author() -> None: + """PymssqlADKMemoryStore binds None for metadata_json=None and missing author.""" + from datetime import datetime, timezone + from typing import cast + + from sqlspec.adapters.pymssql.adk import PymssqlADKMemoryStore + + config = _mock_config() + conn = config.provide_connection.return_value.__enter__.return_value + cursor = conn.cursor.return_value + cursor.rowcount = 1 + store = PymssqlADKMemoryStore(config) + now = datetime.now(tz=timezone.utc) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = store.insert_memory_entries([entry]) + + assert inserted == 1 + params = cursor.execute.call_args.args[1] + assert params[6] is None + assert params[10] is None diff --git a/tests/unit/adapters/test_pymssql/test_core.py b/tests/unit/adapters/test_pymssql/test_core.py index ee1e947a5..b8abaa6be 100644 --- a/tests/unit/adapters/test_pymssql/test_core.py +++ b/tests/unit/adapters/test_pymssql/test_core.py @@ -156,9 +156,15 @@ class AttributeException(Exception): class BoolAttributeException(Exception): number = True + class ZeroAttributeException(Exception): + number = 0 + assert extract_error_number(AttributeException("duplicate key")) == 2627 assert extract_error_number(BoolAttributeException("Msg 2627, Level 14")) == 2627 + assert extract_error_number(ZeroAttributeException("Msg 2627, Level 14")) == 2627 + assert extract_error_number(ZeroAttributeException("Plain error")) is None assert extract_error_number(Exception(True, "Msg 1205, Level 13")) == 1205 + assert extract_error_number(Exception(0, "Msg 1205, Level 13")) == 1205 assert extract_error_number(Exception(1205, "Deadlock found")) == 1205 assert extract_error_number(Exception("Violation of UNIQUE KEY constraint (2627)")) == 2627 assert extract_error_number(Exception("Plain error")) is None diff --git a/tests/unit/adapters/test_pymssql/test_driver.py b/tests/unit/adapters/test_pymssql/test_driver.py index 7a75e550c..f746aee83 100644 --- a/tests/unit/adapters/test_pymssql/test_driver.py +++ b/tests/unit/adapters/test_pymssql/test_driver.py @@ -245,7 +245,7 @@ def fail(*_args: object) -> None: with pytest.raises(SQLSpecError) as caught: getattr(driver, operation)() assert caught.value.__cause__ is failure - assert driver._connection_in_transaction() is (operation != "begin") + assert driver._connection_in_transaction() is (operation == "commit") @pytest.mark.parametrize("fails", [False, True]) diff --git a/tests/unit/adapters/test_pymssql/test_extensions.py b/tests/unit/adapters/test_pymssql/test_extensions.py index c1c328480..6e65c424c 100644 --- a/tests/unit/adapters/test_pymssql/test_extensions.py +++ b/tests/unit/adapters/test_pymssql/test_extensions.py @@ -1,5 +1,6 @@ """pymssql extension package tests.""" +from typing import Any from unittest.mock import AsyncMock import pytest @@ -53,19 +54,35 @@ async def test_litestar_store_async_methods_bridge_sync_operations(monkeypatch: assert calls == ["create"] -def test_adk_store_ddl_uses_tsql_tables_and_json_fallback() -> None: - """ADK DDL should use T-SQL table shape and NVARCHAR JSON fallback by default.""" +@pytest.mark.parametrize(("major", "expected_json_type"), [(16, "NVARCHAR(MAX)"), (17, "JSON")]) +def test_adk_store_ddl_uses_tsql_tables_and_json_fallback( + monkeypatch: pytest.MonkeyPatch, major: int, expected_json_type: str +) -> None: + """ADK DDL should detect SQL Server version lazily when native_json is not configured.""" + from contextlib import contextmanager + from unittest.mock import MagicMock + from sqlspec.adapters.pymssql.adk.store import PymssqlADKStore + from sqlspec.adapters.pymssql.data_dictionary import MssqlVersionInfo + + config = PymssqlConfig(extension_config={"adk": {}}) + driver = MagicMock() + driver.data_dictionary.get_version.return_value = MssqlVersionInfo(major=major) - store = PymssqlADKStore(PymssqlConfig(extension_config={"adk": {}})) + @contextmanager + def _fake_session(*_args: Any, **_kwargs: Any) -> Any: + yield driver + + monkeypatch.setattr(PymssqlConfig, "provide_session", _fake_session) + store = PymssqlADKStore(config) sessions_ddl = store._sessions_table_ddl() events_ddl = store._events_table_ddl() assert "CREATE TABLE" in sessions_ddl - assert "NVARCHAR(MAX)" in sessions_ddl + assert f"state {expected_json_type} NOT NULL" in sessions_ddl assert "SYSUTCDATETIME()" in sessions_ddl - assert "event_data" in events_ddl + assert f"event_data {expected_json_type} NOT NULL" in events_ddl assert "DATETIME2(6)" in events_ddl @@ -76,3 +93,12 @@ def test_adk_store_can_force_native_json_column_type() -> None: store = PymssqlADKStore(PymssqlConfig(extension_config={"adk": {"native_json": True}})) assert "state JSON NOT NULL" in store._sessions_table_ddl() + + +def test_adk_store_can_force_fallback_json_column_type() -> None: + """ADK config should allow forcing NVARCHAR(MAX) without opening a session.""" + from sqlspec.adapters.pymssql.adk.store import PymssqlADKStore + + store = PymssqlADKStore(PymssqlConfig(extension_config={"adk": {"native_json": False}})) + + assert "state NVARCHAR(MAX) NOT NULL" in store._sessions_table_ddl() diff --git a/tests/unit/adapters/test_pymssql/test_pool.py b/tests/unit/adapters/test_pymssql/test_pool.py index 35e3249a6..219101c9e 100644 --- a/tests/unit/adapters/test_pymssql/test_pool.py +++ b/tests/unit/adapters/test_pymssql/test_pool.py @@ -100,9 +100,11 @@ def _open() -> None: worker.join() assert len({id(connection) for connection in opened}) == 2 + assert pool.size() == 2 pool.close() + assert pool.size() == 0 assert [connection.closed for connection in opened] == [True, True] diff --git a/tests/unit/adapters/test_pymysql/test_adk_store.py b/tests/unit/adapters/test_pymysql/test_adk_store.py index 91f2c3fb1..c952e77ea 100644 --- a/tests/unit/adapters/test_pymysql/test_adk_store.py +++ b/tests/unit/adapters/test_pymysql/test_adk_store.py @@ -211,3 +211,93 @@ def test_pymysql_list_sessions_rejects_invalid_options(options: "dict[str, Any]" store.list_sessions("app", **options) assert cursor.calls == [] + + +def test_pymysql_adk_memory_store_insert_and_search_json_handling() -> None: + """PyMysqlADKMemoryStore handles missing author, metadata_json=None, and deserializes JSON strings.""" + from datetime import datetime, timezone + + cursor = MagicMock() + cursor.rowcount = 1 + now = datetime.now(tz=timezone.utc) + cursor.description = [ + ("id",), + ("session_id",), + ("app_name",), + ("user_id",), + ("scope",), + ("event_id",), + ("author",), + ("timestamp",), + ("content_json",), + ("content_text",), + ("metadata_json",), + ("inserted_at",), + ] + cursor.fetchall.return_value = [ + ( + "mem-1", + "sess-1", + "app", + "user-1", + "user", + "evt-1", + None, + now, + '{"text": "hello"}', + "hello", + '{"source": "unit"}', + now, + ) + ] + conn = MagicMock() + conn.cursor.return_value = cursor + conn.__enter__.return_value = conn + conn.__exit__.return_value = None + config = _mock_config({"owner_id_column": "owner_id INT NULL"}) + config.provide_connection = lambda *_a, **_k: conn + store = PyMysqlADKMemoryStore(config) + entry = cast( + "Any", + { + "id": "mem-1", + "session_id": "sess-1", + "app_name": "app", + "user_id": "user-1", + "event_id": "evt-1", + "timestamp": now, + "content_json": {"text": "hello"}, + "content_text": "hello", + "metadata_json": None, + "inserted_at": now, + }, + ) + + inserted = store._insert_memory_entries([entry], owner_id=99) + records = store._search_entries("hello", "app", "user-1") + + assert inserted == 1 + insert_params = cursor.execute.call_args_list[0].args[1] + assert insert_params[6] is None + assert insert_params[11] is None + assert len(records) == 1 + assert records[0]["content_json"] == {"text": "hello"} + assert records[0]["metadata_json"] == {"source": "unit"} + + +def test_pymysql_stream_source_closes_cursor_when_execute_raises() -> None: + """PymysqlStreamSource.start() closes the cursor if execute() fails.""" + from sqlspec.adapters.pymysql.core import PymysqlStreamSource + + cursor = MagicMock() + cursor.execute.side_effect = RuntimeError("execute boom") + connection = MagicMock() + connection.cursor.return_value = cursor + driver = MagicMock(connection=connection) + source = PymysqlStreamSource(driver, "SELECT 1", (), 100, set()) + + with pytest.raises(RuntimeError, match="execute boom"): + source.start() + + cursor.close.assert_called_once() + assert source._cursor is None diff --git a/tests/unit/adapters/test_pymysql/test_cloud_sql_connector.py b/tests/unit/adapters/test_pymysql/test_cloud_sql_connector.py index 541442dc4..4c5431518 100644 --- a/tests/unit/adapters/test_pymysql/test_cloud_sql_connector.py +++ b/tests/unit/adapters/test_pymysql/test_cloud_sql_connector.py @@ -109,7 +109,7 @@ def test_cloud_sql_setup_strips_direct_connection_parameters(mock_cloud_sql_modu pool = config._create_pool() mock_cloud_sql_module.assert_called_once() - assert config.get_cloud_sql_connector() is mock_connector + assert config._get_cloud_sql_connector() is mock_connector assert pool._connection_factory is not None assert "host" not in pool._connection_parameters assert "port" not in pool._connection_parameters @@ -195,4 +195,4 @@ def test_cloud_sql_connector_cleanup(mock_cloud_sql_module: MagicMock) -> None: config._close_pool() mock_connector.close.assert_called_once() - assert config.get_cloud_sql_connector() is None + assert config._get_cloud_sql_connector() is None diff --git a/tests/unit/adapters/test_spanner/test_adk_store.py b/tests/unit/adapters/test_spanner/test_adk_store.py index 3e078156f..c46e1ac42 100644 --- a/tests/unit/adapters/test_spanner/test_adk_store.py +++ b/tests/unit/adapters/test_spanner/test_adk_store.py @@ -17,6 +17,7 @@ SpannerSyncADKMemoryStore, SpannerSyncADKStore, ) +from sqlspec.adapters.spanner.adk.store import _spanner_drop_statement_table from sqlspec.config import ADKConfig from sqlspec.extensions.adk import StoredEvent, StoredMemory @@ -255,6 +256,19 @@ def test_spanner_memory_reset_drop_tables_filters_absent_tables_and_indexes() -> ] +def test_spanner_drop_statement_table_handles_if_exists_and_quoted_identifiers() -> None: + existing = {"adk_events", "adk_memory_entries"} + + assert _spanner_drop_statement_table("DROP TABLE IF EXISTS `adk_events`", existing) == "adk_events" + assert _spanner_drop_statement_table("DROP TABLE IF EXISTS `adk_session`", existing) is None + assert _spanner_drop_statement_table("DROP INDEX IF EXISTS `idx_adk_events_timestamp`", existing) == "adk_events" + assert ( + _spanner_drop_statement_table("DROP SEARCH INDEX IF EXISTS `idx_adk_memory_entries_fts`", existing) + == "adk_memory_entries" + ) + assert _spanner_drop_statement_table("DROP SEARCH INDEX `idx_adk_session_fts`", existing) is None + + def test_get_session_returns_none_when_spanner_session_table_missing() -> None: store = SpannerSyncADKStore(_mock_config()) diff --git a/tests/unit/adapters/test_spanner/test_batch_write_api.py b/tests/unit/adapters/test_spanner/test_batch_write_api.py index cc4c71063..9dca5a161 100644 --- a/tests/unit/adapters/test_spanner/test_batch_write_api.py +++ b/tests/unit/adapters/test_spanner/test_batch_write_api.py @@ -114,7 +114,7 @@ def test_batch_write_overwrite_uses_transactional_mutations(batch_write_driver: conn = cast("_FakeBatchTransaction", batch_write_driver.connection) batch_write_driver.load_from_arrow("users", pa.table({"id": [1]}), overwrite=True) - assert conn.execute_update_calls and "DELETE FROM users WHERE TRUE" in conn.execute_update_calls[0] + assert conn.execute_update_calls and "DELETE FROM `users` WHERE TRUE" in conn.execute_update_calls[0] assert conn.insert_or_update_calls == [("users", ["id"], [[1]])] assert conn.database.mutation_groups_obj.batch_write_calls == 0 diff --git a/tests/unit/adapters/test_spanner/test_config.py b/tests/unit/adapters/test_spanner/test_config.py index 3fa465353..39f470588 100644 --- a/tests/unit/adapters/test_spanner/test_config.py +++ b/tests/unit/adapters/test_spanner/test_config.py @@ -169,7 +169,10 @@ def instance(self, instance_id: str, **kwargs: Any) -> _FakeInstance: created_clients: list[_FakeClient] = [] import google.cloud.spanner_v1 + import sqlspec.adapters.spanner.config as spanner_config + monkeypatch.setattr(google.cloud.spanner_v1, "Client", _FakeClient) + monkeypatch.setattr(spanner_config, "Client", _FakeClient) client_info = object() query_options = object() diff --git a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py index f459d5848..c562f0102 100644 --- a/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py +++ b/tests/unit/adapters/test_spanner/test_load_from_arrow_mutations.py @@ -71,10 +71,39 @@ def test_load_from_arrow_overwrite_deletes_then_mutates(mutations_driver: Spanne mutations_driver.load_from_arrow("users", arrow_table, overwrite=True) assert txn.execute_update_calls - assert "DELETE FROM users WHERE TRUE" in txn.execute_update_calls[0] + assert "DELETE FROM `users` WHERE TRUE" in txn.execute_update_calls[0] assert len(txn.insert_or_update_calls) == 1 +def test_load_from_arrow_overwrite_quotes_qualified_table_identifier(mutations_driver: SpannerSyncDriver) -> None: + txn = cast("_FakeTransaction", mutations_driver.connection) + arrow_table = pa.table({"id": [1]}) + + mutations_driver.load_from_arrow("my_schema.order", arrow_table, overwrite=True) + + assert txn.execute_update_calls == ["DELETE FROM `my_schema`.`order` WHERE TRUE"] + + +def test_dispatch_execute_script_recognizes_union_and_spanner_queries_as_reads() -> None: + class _ReadSnapshot: + def __init__(self) -> None: + self.executed_sql: list[str] = [] + + def execute_sql(self, sql: str, **_kwargs: Any) -> list[Any]: + self.executed_sql.append(sql) + return [] + + snapshot = _ReadSnapshot() + driver = SpannerSyncDriver(cast("Any", snapshot)) + script = ( + "SELECT 1 UNION ALL SELECT 2; @{FORCE_INDEX=_BASE_TABLE} SELECT * FROM users @{FORCE_INDEX=idx_users_name};" + ) + result = driver.dispatch_execute_script(cast("Any", snapshot), driver.prepare_statement(script)) + + assert result.statement_count == 2 + assert len(snapshot.executed_sql) == 2 + + def test_load_from_arrow_rerun_is_idempotent(mutations_driver: SpannerSyncDriver) -> None: txn = cast("_FakeTransaction", mutations_driver.connection) arrow_table = pa.table({"id": [1, 2], "name": ["a", "b"]}) diff --git a/tests/unit/builder/test_merge.py b/tests/unit/builder/test_merge.py index cad2e37a9..200cff3fa 100644 --- a/tests/unit/builder/test_merge.py +++ b/tests/unit/builder/test_merge.py @@ -197,6 +197,29 @@ def test_merge_when_matched_update_with_condition_oracle() -> None: assert "UPDATE SET price = src.new_price WHERE src.new_price < t.price" in rendered +def test_merge_oracle_build_does_not_mutate_subsequent_postgres_build() -> None: + """Building a MERGE query for Oracle must not mutate the AST for subsequent Postgres builds.""" + query = ( + sql + .merge() + .into("products", alias="t") + .using("staging", alias="src") + .on("t.id = src.id") + .when_matched_then_update({"price": "src.new_price"}, condition="src.new_price < t.price") + ) + + oracle_stmt = query.build(dialect="oracle") + postgres_stmt = query.build(dialect="postgres") + + oracle_sql = " ".join(oracle_stmt.sql.split()) + postgres_sql = " ".join(postgres_stmt.sql.split()) + assert "UPDATE SET price = src.new_price WHERE src.new_price < t.price" in oracle_sql + assert ( + 'WHEN MATCHED AND "src"."new_price" < "t"."price" THEN UPDATE SET "price" = "src"."new_price"' in postgres_sql + ) + assert "WHERE" not in postgres_sql + + def test_merge_when_matched_update_no_values_error() -> None: """Test that WHEN MATCHED UPDATE without values raises error.""" query = sql.merge().into("products", alias="t").using("staging", alias="s").on("t.id = s.id") diff --git a/tests/unit/dialects/test_spanner_hints.py b/tests/unit/dialects/test_spanner_hints.py index 14e0f7226..61d15f541 100644 --- a/tests/unit/dialects/test_spanner_hints.py +++ b/tests/unit/dialects/test_spanner_hints.py @@ -109,3 +109,16 @@ def test_join_hint_round_trip() -> None: def test_at_comment_on_expression_is_preserved() -> None: sql = "SELECT 1 /* @ ordinary comment */" assert "@ ordinary comment" in parse_one(sql, dialect="spanner").sql(dialect="spanner") + + +def test_at_comments_without_equals_on_select_table_and_join_are_not_hints() -> None: + sql = "-- @author: Alice\nSELECT * FROM Albums /* @todo */ JOIN /* @note */ Singers ON Albums.SingerId = Singers.Id" + parsed = parse_one(sql, dialect="spanner") + + assert parsed.args.get("hint") is None + table = parsed.find(exp.Table) + assert table is not None + assert not table.args.get("hints") + join = parsed.find(exp.Join) + assert join is not None + assert join.args.get("spanner_hint") is None diff --git a/tests/unit/extensions/test_adk/test_memory_converters.py b/tests/unit/extensions/test_adk/test_memory_converters.py index e32cda6d5..85e7b493a 100644 --- a/tests/unit/extensions/test_adk/test_memory_converters.py +++ b/tests/unit/extensions/test_adk/test_memory_converters.py @@ -10,15 +10,34 @@ from google.adk.events.event import Event from google.adk.events.event_actions import EventActions +from google.adk.memory.memory_entry import MemoryEntry from google.adk.sessions.session import Session from google.genai import types +from sqlspec.extensions.adk.memory._types import StoredMemory from sqlspec.extensions.adk.memory.converters import ( event_to_memory_record, extract_content_text, + memory_entry_to_record, record_to_memory_entry, session_to_memory_records, ) +from sqlspec.extensions.adk.memory.service import SQLSpecSyncMemoryService + + +class _MockSyncMemoryStore: + def __init__(self) -> None: + self.enabled = True + self.inserted_batches: list[list[StoredMemory]] = [] + + def insert_memory_entries(self, entries: list[StoredMemory], owner_id: object | None = None) -> int: + self.inserted_batches.append(entries) + return len(entries) + + def search_entries( + self, *, query: str, app_name: str, user_id: str, limit: int | None = None + ) -> list[StoredMemory]: + return [] def _event(event_id: str, text: str | None) -> Event: @@ -67,3 +86,43 @@ def test_session_to_memory_records_roundtrip() -> None: assert entry.content is not None assert entry.content.parts is not None assert entry.content.parts[0].text == "Hello memory" + + +def test_memory_entry_to_record_generates_unique_event_id_when_missing() -> None: + entry_a = MemoryEntry(content=types.Content(parts=[types.Part(text="First memory")])) + entry_b = MemoryEntry(content=types.Content(parts=[types.Part(text="Second memory")])) + + record_a = memory_entry_to_record(entry_a, app_name="app", user_id="user") + record_b = memory_entry_to_record(entry_b, app_name="app", user_id="user") + + assert record_a is not None + assert record_b is not None + assert record_a["event_id"] == record_a["id"] + assert record_b["event_id"] == record_b["id"] + assert record_a["event_id"] != "" + assert record_a["event_id"] != record_b["event_id"] + assert record_a["session_id"] == "__unknown_session_id__" + + +def test_sync_memory_service_add_events_and_add_memory() -> None: + store = _MockSyncMemoryStore() + service = SQLSpecSyncMemoryService(store) # type: ignore[arg-type] + + service.add_events_to_memory( + app_name="app", user_id="user", events=[_event("evt-1", "Incremental memory")], session_id=None + ) + service.add_memory( + app_name="app", + user_id="user", + memories=[ + MemoryEntry(content=types.Content(parts=[types.Part(text="Direct memory 1")])), + MemoryEntry(content=types.Content(parts=[types.Part(text="Direct memory 2")])), + ], + ) + + assert len(store.inserted_batches) == 2 + assert store.inserted_batches[0][0]["session_id"] == "__unknown_session_id__" + assert store.inserted_batches[0][0]["event_id"] == "evt-1" + assert len(store.inserted_batches[1]) == 2 + assert store.inserted_batches[1][0]["event_id"] != store.inserted_batches[1][1]["event_id"] + assert store.inserted_batches[1][0]["session_id"] == "__unknown_session_id__" diff --git a/tests/unit/extensions/test_adk/test_service.py b/tests/unit/extensions/test_adk/test_service.py index 4d3b67aa5..c60978e2c 100644 --- a/tests/unit/extensions/test_adk/test_service.py +++ b/tests/unit/extensions/test_adk/test_service.py @@ -19,6 +19,8 @@ if importlib.util.find_spec("google.genai") is None or importlib.util.find_spec("google.adk") is None: pytest.skip("google-adk not installed", allow_module_level=True) +from google.adk.errors import StaleSessionError +from google.adk.errors.session_not_found_error import SessionNotFoundError from google.adk.events.event import Event from google.adk.events.event_actions import EventActions from google.adk.sessions.base_session_service import GetSessionConfig @@ -612,18 +614,17 @@ async def test_get_user_state_strips_user_prefixes() -> None: @pytest.mark.anyio async def test_append_event_raises_on_stale_marker() -> None: - """append_event raises ValueError when _storage_update_marker doesn't match storage.""" + """append_event raises StaleSessionError when _storage_update_marker doesn't match storage.""" from sqlspec.extensions.adk.converters import compute_update_marker store = StaleDetectionStore(stale_marker=True) service = SQLSpecSessionService(store) # type: ignore[arg-type] session = _make_session() - # Set a marker that won't match the advanced update_time session._storage_update_marker = compute_update_marker(store._session_record["update_time"]) # type: ignore[arg-type] event = _make_event() - with pytest.raises(ValueError, match="modified in storage"): + with pytest.raises(StaleSessionError, match="modified in storage"): await service.append_event(session, event) assert not store.append_event_and_update_state_called @@ -631,16 +632,15 @@ async def test_append_event_raises_on_stale_marker() -> None: @pytest.mark.anyio async def test_append_event_raises_on_stale_timestamp() -> None: - """append_event raises ValueError when storage update_time > session.last_update_time.""" + """append_event raises StaleSessionError when storage update_time > session.last_update_time.""" store = StaleDetectionStore(stale_timestamp=True) service = SQLSpecSessionService(store) # type: ignore[arg-type] - # Session loaded with an older timestamp; marker is None (timestamp fallback) session = _make_session() - session._storage_update_marker = None # force timestamp-based check + session._storage_update_marker = None event = _make_event() - with pytest.raises(ValueError, match="modified in storage"): + with pytest.raises(StaleSessionError, match="modified in storage"): await service.append_event(session, event) assert not store.append_event_and_update_state_called @@ -648,13 +648,13 @@ async def test_append_event_raises_on_stale_timestamp() -> None: @pytest.mark.anyio async def test_append_event_raises_when_session_not_found() -> None: - """append_event raises ValueError when the session no longer exists in storage.""" + """append_event raises SessionNotFoundError when the session no longer exists in storage.""" store = MissingSessionStore() service = SQLSpecSessionService(store) # type: ignore[arg-type] session = _make_session() event = _make_event() - with pytest.raises(ValueError, match="not found"): + with pytest.raises(SessionNotFoundError, match="not found"): await service.append_event(session, event) assert not store.append_event_and_update_state_called diff --git a/tests/unit/migrations/test_schema_ensure.py b/tests/unit/migrations/test_schema_ensure.py index d79a58632..fd2e65ac3 100644 --- a/tests/unit/migrations/test_schema_ensure.py +++ b/tests/unit/migrations/test_schema_ensure.py @@ -139,3 +139,21 @@ async def test_async_ensure_diff_adds_missing_column() -> None: assert result.added_columns == {"widgets": ["label"]} driver.execute.assert_awaited_once() driver.commit.assert_awaited_once() + + +def test_schema_target_from_ddl_ignores_parentheses_in_strings_and_comments() -> None: + ddl = """ + BEGIN + EXECUTE IMMEDIATE 'CREATE TABLE widgets ( + id NUMBER PRIMARY KEY, + /* comment with ) inside */ + delim VARCHAR2(10) DEFAULT '')'', + label VARCHAR2(50) NOT NULL + )'; + END; + """ + + target = SchemaTarget.from_ddl("widgets", ddl, dialect="oracle") + + assert [column.name for column in target.create_table.columns] == ["id", "delim", "label"] + assert target.create_table.columns[2].not_null is True diff --git a/tools/scripts/mypyc_inventory.py b/tools/scripts/mypyc_inventory.py index e618f1cdd..7cde43ed2 100644 --- a/tools/scripts/mypyc_inventory.py +++ b/tools/scripts/mypyc_inventory.py @@ -10,6 +10,7 @@ __all__ = ( "HOT_SURFACE_CLASSIFICATIONS", + "PRESERVED_EXCLUSIONS", "build_inventory", "classify_module", "classify_surface", @@ -17,6 +18,7 @@ "list_sqlspec_modules", "load_mypyc_patterns", "main", + "validate_inventory", ) try: @@ -73,35 +75,47 @@ "classification": "compile_now", "reason": "Uses importlib.resources instead of direct __file__ path discovery.", }, - "sqlspec/data_dictionary/dialects/postgres.py": { + "sqlspec/data_dictionary/dialects/postgres/config.py": { "classification": "compile_now", "reason": "Shared Postgres JSON type helper for ADBC-as-Postgres, asyncpg, psqlpy, and psycopg dictionaries.", }, - "sqlspec/data_dictionary/dialects/sqlite.py": { + "sqlspec/data_dictionary/dialects/sqlite/config.py": { "classification": "compile_now", "reason": "Shared SQLite JSON and feature-list helpers for sqlite, aiosqlite, and ADBC-as-SQLite dictionaries.", }, - "sqlspec/data_dictionary/dialects/mysql.py": { + "sqlspec/data_dictionary/dialects/mysql/config.py": { "classification": "compile_now", "reason": "Shared MySQL JSON type helper for mysqlconnector, pymysql, aiomysql, asyncmy, and ADBC-as-MySQL dictionaries.", }, - "sqlspec/data_dictionary/dialects/cockroachdb.py": { + "sqlspec/data_dictionary/dialects/mysql/dictionary.py": { + "classification": "hard_block", + "reason": "Cross-module data-dictionary inheritance stays interpreted to avoid mypyc segfaults.", + }, + "sqlspec/data_dictionary/dialects/cockroachdb/config.py": { "classification": "compile_now", "reason": "Shared CockroachDB JSON type helper for cockroach_asyncpg, cockroach_psycopg, and ADBC-as-Cockroach dictionaries.", }, - "sqlspec/data_dictionary/dialects/duckdb.py": { + "sqlspec/data_dictionary/dialects/db2/config.py": { + "classification": "compile_now", + "reason": "Db2 data-dictionary dialect configuration and feature helpers compile with the dialect surface.", + }, + "sqlspec/data_dictionary/dialects/duckdb/config.py": { "classification": "compile_now", "reason": "DuckDB data-dictionary dialect configuration is in the compiled dialect surface.", }, - "sqlspec/data_dictionary/dialects/oracle.py": { + "sqlspec/data_dictionary/dialects/mssql/config.py": { + "classification": "compile_now", + "reason": "Shared SQL Server version, feature, and JSON helpers compile with the dialect surface.", + }, + "sqlspec/data_dictionary/dialects/oracle/config.py": { "classification": "compile_now", "reason": "Shared Oracle version, JSON, feature, and table-list helpers for oracledb sync and async dictionaries.", }, - "sqlspec/data_dictionary/dialects/spanner.py": { + "sqlspec/data_dictionary/dialects/spanner/config.py": { "classification": "compile_now", "reason": "Spanner data-dictionary dialect configuration is in the compiled dialect surface.", }, - "sqlspec/data_dictionary/dialects/bigquery.py": { + "sqlspec/data_dictionary/dialects/bigquery/config.py": { "classification": "compile_now", "reason": "Shared BigQuery INFORMATION_SCHEMA formatting helpers for native BigQuery and ADBC-as-BigQuery dictionaries.", }, @@ -125,10 +139,18 @@ "classification": "hard_block", "reason": "SQLGlot tokenizer/dialect subclass module fails native class import under mypyc; compiled helpers stay in _generators/_operators.", }, + "sqlspec/dialects/spanner/_expressions.py": { + "classification": "keep_interpreted", + "reason": "Spanner SQLGlot AST expression builders remain interpreted alongside the dialect subclasses.", + }, "sqlspec/dialects/spanner/_generators.py": { "classification": "compile_now", "reason": "Spanner SQL rendering helpers compile with the custom Spanner dialect surface.", }, + "sqlspec/dialects/spanner/_parsers.py": { + "classification": "compile_now", + "reason": "Spanner PROPERTY_PARSERS entries compile as explicit-argument callables.", + }, "sqlspec/dialects/spanner/_spangres.py": { "classification": "hard_block", "reason": "SQLGlot subclass/registration module fails native class import under mypyc; compiled helpers stay in _generators.", @@ -143,6 +165,10 @@ }, "sqlspec/dialects/db2/_parsers.py": {"classification": "compile_now", "reason": "Db2 AST normalisation helpers."}, "sqlspec/dialects/db2/_transforms.py": {"classification": "compile_now", "reason": "Db2 render helper functions."}, + "sqlspec/adapters/mysql_common.py": { + "classification": "compile_now", + "reason": "Shared MySQL-family adapter helpers compile with the adapter core surface.", + }, "sqlspec/extensions/events/_models.py": { "classification": "compile_now", "reason": "EventMessage has concrete datetime annotations and slot dataclass layout compatible with mypyc.", @@ -295,6 +321,24 @@ }, } +PRESERVED_EXCLUSIONS: frozenset[str] = frozenset({ + "sqlspec/dialects/postgres/_paradedb.py", + "sqlspec/dialects/postgres/_pg_textsearch.py", + "sqlspec/dialects/postgres/_pgvector.py", + "sqlspec/dialects/spanner/_expressions.py", + "sqlspec/dialects/spanner/_spangres.py", + "sqlspec/dialects/spanner/_spanner.py", + "sqlspec/utils/arrow_helpers.py", + "sqlspec/storage/_arrow_payload.py", + "sqlspec/adapters/**/data_dictionary.py", + "sqlspec/data_dictionary/dialects/mysql/dictionary.py", + "sqlspec/migrations/commands.py", + "sqlspec/extensions/events/_store.py", + "sqlspec/extensions/adk/converters.py", + "sqlspec/config.py", + "sqlspec/core/_pagination.py", +}) + def load_mypyc_patterns(root: Path) -> tuple[list[str], list[str]]: """Load mypyc include/exclude glob patterns from pyproject.toml.""" @@ -419,33 +463,37 @@ def build_inventory(root: Path | None = None) -> dict[str, Any]: "status": "compiled", "classification": "compile_now", }, - "preserved_exclusions": sorted( - pattern - for pattern in exclude_patterns - if pattern - in { - "sqlspec/dialects/postgres/_paradedb.py", - "sqlspec/dialects/postgres/_pg_textsearch.py", - "sqlspec/dialects/postgres/_pgvector.py", - "sqlspec/dialects/spanner/_spangres.py", - "sqlspec/dialects/spanner/_spanner.py", - "sqlspec/utils/arrow_helpers.py", - "sqlspec/storage/_arrow_payload.py", - "sqlspec/adapters/**/data_dictionary.py", - "sqlspec/observability/_formatting.py", - "sqlspec/migrations/commands.py", - "sqlspec/extensions/events/_channel.py", - "sqlspec/extensions/events/_models.py", - "sqlspec/extensions/events/_queue.py", - "sqlspec/extensions/events/_store.py", - "sqlspec/extensions/adk/converters.py", - "sqlspec/config.py", - } - ), + "preserved_exclusions": sorted(pattern for pattern in exclude_patterns if pattern in PRESERVED_EXCLUSIONS), "hot_surfaces": hot_surfaces, } +def validate_inventory(root: Path | None = None) -> list[str]: + """Validate that hot-surface classifications and preserved exclusions match pyproject.toml.""" + project_root = root or Path(__file__).resolve().parents[2] + include_patterns, exclude_patterns = load_mypyc_patterns(project_root) + module_set = set(list_sqlspec_modules(project_root)) + exclude_set = set(exclude_patterns) + errors: list[str] = [] + + for module_path, details in sorted(HOT_SURFACE_CLASSIFICATIONS.items()): + if module_path not in module_set: + errors.append(f"hot surface module does not exist: {module_path}") + continue + status = classify_module(module_path, include_patterns, exclude_patterns) + classification = details["classification"] + if classification == "compile_now" and status != "compiled": + errors.append(f"hot surface marked compile_now is not compiled: {module_path}") + elif classification != "compile_now" and status == "compiled": + errors.append(f"hot surface marked {classification} is compiled: {module_path}") + + errors.extend( + f"preserved exclusion missing from pyproject.toml: {pattern}" + for pattern in sorted(PRESERVED_EXCLUSIONS - exclude_set) + ) + return errors + + def format_markdown(inventory: dict[str, Any]) -> str: """Format inventory output as markdown.""" summary = inventory["summary"] @@ -495,8 +543,18 @@ def main(argv: Sequence[str] | None = None) -> int: default=Path(__file__).resolve().parents[2], help="Project root containing pyproject.toml and sqlspec/.", ) + parser.add_argument( + "--check", action="store_true", help="Verify that hot surfaces and preserved exclusions match pyproject.toml." + ) args = parser.parse_args(argv) + if args.check: + errors = validate_inventory(args.root) + if errors: + for error in errors: + sys.stderr.write(f"{error}\n") + return 1 + inventory = build_inventory(args.root) if args.format == "markdown": sys.stdout.write(format_markdown(inventory)) diff --git a/tools/scripts/mypyc_smoke.py b/tools/scripts/mypyc_smoke.py index d95628a2b..c53fe936a 100644 --- a/tools/scripts/mypyc_smoke.py +++ b/tools/scripts/mypyc_smoke.py @@ -577,6 +577,8 @@ def _discover_adapter_config_classes(*, skipped: "list[str] | None" = None) -> " "asyncpg", "duckdb", "google", + "ibm_db", + "ibm_db_dbi", "mssql_python", "mysql", "oracledb", From e36ccf3f924e8cefd77db610bc598b6ff66cb863 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Tue, 29 Sep 2026 00:13:35 +0000 Subject: [PATCH 15/15] fix: preserve runtime adapter types and dialect-aware DDL parsing --- docs/changelog.rst | 23 +++++-- docs/reference/adapters/duckdb.rst | 10 +++ docs/reference/adapters/spanner.rst | 9 +++ sqlspec/adapters/db2/_typing.py | 8 +-- sqlspec/adapters/duckdb/config.py | 65 +++++++------------ sqlspec/adapters/duckdb/driver.py | 5 +- sqlspec/adapters/mysqlconnector/_typing.py | 4 +- sqlspec/migrations/schema.py | 37 +++++++---- tests/unit/adapters/test_db2/test_config.py | 13 ++++ .../test_duckdb/test_arrow_streaming.py | 28 ++++---- .../test_mysqlconnector/test_config.py | 13 ++++ tests/unit/migrations/test_schema_ensure.py | 37 +++++++++++ 12 files changed, 170 insertions(+), 82 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a6bfdcff2..e57605ee8 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -27,13 +27,11 @@ Unreleased and later when the SQLite runtime supports them. Choose a transaction lock mode or set the batch size for Arrow imports. Defaults stay the same. (`#820 `_) -* Spanner forwards native query options (``QueryOptions``, ``RequestOptions``, - ``DirectedReadOptions``) and final-statement hints, supports opt-in Batch Write - from read sessions, adds ``run_in_transaction`` retry with backoff on ``Aborted`` - and ``last_statement=True`` commit inlining, supports multiplexed session pooling - and ``PingingPool`` keepalive intervals, coerces ``FLOAT32`` vector parameters, - and adds ``execute_partitioned_dml()`` for Partitioned DML truncation during - ``load_from_arrow(overwrite=True)``. +* Spanner forwards native query, request, and directed-read options through + existing execution and session APIs. ``last_statement=True`` marks the final + DML statement; it does not commit the transaction. Opt-in Batch Write accepts + read sessions for Arrow imports without overwrite. Overwrite uses a + transaction for both the delete and replacement mutations. (`#814 `_) * Cloud Spanner and Spangres SQLGlot dialects isolate custom ``SpannerParser``, ``SpangresParser``, ``SpannerGenerator``, and ``SpangresGenerator`` subclasses @@ -90,6 +88,17 @@ Unreleased **Fixed:** +* MySQL Connector can create async pools again. Db2 exposes the installed + driver's connection and cursor types to apps. + +* DuckDB returns UUID objects for UUID columns on the first query, on cache + hits, and in row streams. Disabling UUID input conversion does not change + result types. Text columns still return strings. + +* Schema checks can read table DDL inside SQL blocks with dialect-specific + quotes. They also handle Oracle blocks whose trailing text is not supported + by the parser. + * Spanner binds Decimal values as NUMERIC and boolean arrays as BOOL. Typed null dictionaries use JSON, and JSON null results stay ``None``. (`#814 `_) diff --git a/docs/reference/adapters/duckdb.rst b/docs/reference/adapters/duckdb.rst index 9db2004b4..acbfe3340 100644 --- a/docs/reference/adapters/duckdb.rst +++ b/docs/reference/adapters/duckdb.rst @@ -6,6 +6,16 @@ Sync DuckDB adapter with full Arrow integration, extension management, and secret configuration. DuckDB excels at analytical workloads and can query Parquet, CSV, and JSON files directly. +UUID parameters and results +=========================== + +``driver_features={"enable_uuid_conversion": False}`` passes UUID parameters +to DuckDB unchanged. The default converts these parameters to strings. +This setting does not change result types: UUID columns return Python UUID +objects, including on cache hits and in row streams. VARCHAR columns remain +strings, even when their values look like UUIDs. Arrow results keep their +native Arrow representation. + Native object-store transfers ============================= diff --git a/docs/reference/adapters/spanner.rst b/docs/reference/adapters/spanner.rst index b95400ba1..28f7403c6 100644 --- a/docs/reference/adapters/spanner.rst +++ b/docs/reference/adapters/spanner.rst @@ -27,6 +27,10 @@ Default request controls can be configured through ``retry`` and ``timeout`` Forwarded to Spanner statement execution calls when provided. +``query_options`` + Forwarded to ``execute_sql()`` and ``execute_update()``. Batch DML does not + accept query options. + Per-call overrides use the existing ``execute()``, ``execute_many()``, and ``execute_script()`` methods: @@ -45,6 +49,11 @@ the argument for a DML statement so call sites can share option plumbing, but it does not forward directed-read options to ``execute_update()`` or ``batch_update()``. +Pass ``last_statement=True`` to mark the final DML request in a transaction. +For scripts, SQLSpec forwards it only when the final statement is DML. This +option does not commit the transaction; commit through the normal transaction +context or driver API. + Session-Scoped Controls ======================= diff --git a/sqlspec/adapters/db2/_typing.py b/sqlspec/adapters/db2/_typing.py index 9c7f3d458..a9323476a 100644 --- a/sqlspec/adapters/db2/_typing.py +++ b/sqlspec/adapters/db2/_typing.py @@ -26,10 +26,10 @@ class _Db2UnavailableError(Exception): from sqlspec.core import StatementConfig if not TYPE_CHECKING: - Db2SyncConnection = Any - Db2RawCursor = Any - Db2AsyncConnection = Any - Db2AsyncRawCursor = Any + Db2SyncConnection = import_optional_attr("ibm_db_dbi", "Connection") or Any + Db2RawCursor = import_optional_attr("ibm_db_dbi", "Cursor") or Any + Db2AsyncConnection = import_optional_attr("ibm_db_dbi", "AsyncConnection") or Any + Db2AsyncRawCursor = import_optional_attr("ibm_db_dbi", "AsyncCursor") or Any Db2Error = import_optional_attr("ibm_db_dbi", "Error") or _Db2UnavailableError ibm_db = import_optional("ibm_db") diff --git a/sqlspec/adapters/duckdb/config.py b/sqlspec/adapters/duckdb/config.py index d1e948ffe..ae1b23264 100644 --- a/sqlspec/adapters/duckdb/config.py +++ b/sqlspec/adapters/duckdb/config.py @@ -165,25 +165,19 @@ class DuckDBDriverFeatures(TypedDict): """TypedDict for DuckDB driver features configuration. Attributes: - extensions: List of extensions to install/load on connection creation. - secrets: List of secrets to create for AI/API integrations. - on_connection_create: Callback executed when connection is created. - json_serializer: Custom JSON serializer for dict/list parameter conversion. - Defaults to sqlspec.utils.serializers.to_json if not provided. - enable_uuid_conversion: Enable automatic UUID string conversion. - When True (default), UUID strings are automatically converted to UUID objects. - When False, UUID strings are treated as regular strings. - extension_flags: Connection-level flags folded into the database startup - configuration. DuckDB rejects these settings once the database is - running, so they cannot be applied afterwards. - enable_events: Enable database event channel support. - Defaults to True when extension_config["events"] is configured. - Provides pub/sub capabilities via table-backed queue (DuckDB has no native pub/sub). - Requires extension_config["events"] for migration setup. - events_backend: Event channel backend selection. - Only option: "poll_queue" (durable table-backed queue with lease-based retries and acknowledgements). - DuckDB does not have native pub/sub, so poll_queue is the only backend. - Defaults to "poll_queue". + extensions: Extensions to install/load when connecting. + secrets: Secrets to create for external services. + on_connection_create: Callback run when connecting. + json_serializer: Serializer for dict/list parameters. Defaults to + sqlspec.utils.serializers.to_json. + enable_uuid_conversion: Convert UUID parameters to strings (default True). + False passes parameters unchanged. UUID result columns always return + UUID objects; VARCHAR columns remain strings. + extension_flags: Startup-only flags; DuckDB rejects changes after startup. + enable_events: Enable table-backed event queues. Defaults to True when + extension_config["events"] is set; requires it for migrations. + events_backend: Only "poll_queue" is supported, with lease-based retries + and acknowledgements. DuckDB has no native pub/sub. """ extensions: NotRequired[Sequence[DuckDBExtensionConfig]] @@ -209,29 +203,16 @@ class _DuckDBSessionConnectionHandler(SyncPoolSessionFactory): class DuckDBConfig(SyncDatabaseConfig[DuckDBConnection, DuckDBConnectionPool, DuckDBDriver]): """DuckDB configuration with connection pooling. - This configuration supports DuckDB's features including: - - - Connection pooling - - Extension management and installation - - Secret management for API integrations - - Auto configuration settings - - Arrow integration - - Direct file querying capabilities - - Configurable type handlers for JSON serialization and UUID conversion - - DuckDB Connection Pool Configuration: - - Default pool size is 1-4 connections (DuckDB uses single connection by default) - - Connection recycling is set to 24 hours by default (set to 0 to disable) - - Shared memory databases use `:memory:shared_db` for proper concurrency - - Type Handler Configuration via driver_features: - - `json_serializer`: Custom JSON serializer for dict/list parameters. - Defaults to `sqlspec.utils.serializers.to_json` if not provided. - Accepts serializer callables that return text. - - - `enable_uuid_conversion`: Enable automatic UUID string conversion (default: True). - When True, UUID strings in query results are automatically converted to UUID objects. - When False, UUID strings are treated as regular strings. + Supports extensions, secrets, Arrow transfers, and direct file queries. + Connections recycle after 24 hours by default; set recycling to 0 to disable. + Shared memory databases use ``:memory:shared_db``. + + Set ``driver_features["json_serializer"]`` to a serializer that returns text + for dict/list parameters. The default is ``sqlspec.utils.serializers.to_json``. + + ``driver_features["enable_uuid_conversion"]`` converts UUID parameters to + strings by default. Disable it to pass UUID parameters through unchanged. + UUID result columns always return UUID objects; VARCHAR columns remain strings. """ driver_type: "ClassVar[type[DuckDBDriver]]" = DuckDBDriver diff --git a/sqlspec/adapters/duckdb/driver.py b/sqlspec/adapters/duckdb/driver.py index 9eeeae7f6..e16888c45 100644 --- a/sqlspec/adapters/duckdb/driver.py +++ b/sqlspec/adapters/duckdb/driver.py @@ -151,8 +151,7 @@ def dispatch_execute(self, cursor: "DuckDBConnection", statement: SQL) -> "Execu if is_select_like: arrow_table = cursor.to_arrow_table() data = arrow_table.to_pylist() - if self.driver_features.get("enable_uuid_conversion", True): - _restore_uuid_columns(data, cursor.description) + _restore_uuid_columns(data, cursor.description) column_names = list(arrow_table.column_names) return self.create_execution_result( @@ -692,7 +691,7 @@ def _open_stream_reader(self, sql: str, parameters: Any, chunk_size: int) -> "tu reader: Any | None = None with handler: result = self.connection.execute(sql, normalize_execute_parameters(parameters)) - description = result.description if self.driver_features.get("enable_uuid_conversion", True) else None + description = result.description reader = result.to_arrow_reader(chunk_size) self._check_pending_exception(handler) if reader is None: diff --git a/sqlspec/adapters/mysqlconnector/_typing.py b/sqlspec/adapters/mysqlconnector/_typing.py index ad74c7fe8..1ad16a74d 100644 --- a/sqlspec/adapters/mysqlconnector/_typing.py +++ b/sqlspec/adapters/mysqlconnector/_typing.py @@ -17,6 +17,8 @@ from mysql.connector.constants import FieldType as _MysqlConnectorFieldType from mysql.connector.cursor import MySQLCursor as _MysqlConnectorSyncRawCursor +from sqlspec.typing import import_optional_attr + if TYPE_CHECKING: from collections.abc import Awaitable, Callable from types import TracebackType @@ -61,7 +63,7 @@ class MysqlConnectorMysqlModuleProtocol(Protocol): MysqlConnectorAsyncRawCursor: TypeAlias = _MysqlConnectorAsyncRawCursor if not TYPE_CHECKING: - MysqlConnectorAsyncPool = Any + MysqlConnectorAsyncPool = import_optional_attr("mysql.connector.aio.pooling", "MySQLConnectionPool") MysqlConnectorAio = _mysql_connector_aio MysqlConnectorSyncConnection = _MysqlConnectorSyncConnection MysqlConnectorAsyncConnection = _MysqlConnectorAsyncConnection diff --git a/sqlspec/migrations/schema.py b/sqlspec/migrations/schema.py index 288113b47..11274416f 100644 --- a/sqlspec/migrations/schema.py +++ b/sqlspec/migrations/schema.py @@ -3,13 +3,17 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any -from sqlglot import TokenType, exp, parse, tokenize +from sqlglot import Dialect, TokenType, exp, parse +from sqlglot.errors import ParseError, TokenError from sqlspec.builder._ddl import AlterTable, CreateTable, _parse_ddl_identifier, _parse_ddl_table if TYPE_CHECKING: from collections.abc import Awaitable, Callable + from sqlglot.dialects.dialect import DialectType + from sqlglot.tokenizer_core import Token + __all__ = ("SchemaEnsureResult", "SchemaTarget", "ensure_schema_async", "ensure_schema_sync") @@ -375,7 +379,7 @@ def _parse_create_expressions(create_statement: str, dialect: Any) -> list[exp.C """Parse direct or procedural-wrapper CREATE TABLE DDL.""" try: expressions = parse(create_statement, read=dialect) - except Exception: + except (ParseError, TokenError): expressions = [] creates = [ expression @@ -384,17 +388,17 @@ def _parse_create_expressions(create_statement: str, dialect: Any) -> list[exp.C ] if creates: return creates - extracted = _extract_create_table_statement(create_statement) + extracted = _extract_create_table_statement(create_statement, dialect) if extracted is None: return [] try: expression = parse(extracted, read=dialect)[0] - except Exception: + except (ParseError, TokenError): return [] return [expression] if isinstance(expression, exp.Create) else [] -def _extract_create_table_statement(create_statement: str) -> "str | None": +def _extract_create_table_statement(create_statement: str, dialect: "DialectType") -> "str | None": """Extract the first balanced CREATE TABLE statement from a wrapper.""" upper_statement = create_statement.upper() start = upper_statement.find("CREATE TABLE") @@ -403,19 +407,13 @@ def _extract_create_table_statement(create_statement: str) -> "str | None": prefix = create_statement[:start].rstrip() if prefix.endswith("'"): quote_pos = create_statement.rfind("'", 0, start) - try: - string_tokens = tokenize(create_statement[quote_pos:]) - except Exception: - return None + string_tokens = _tokenize_ddl_prefix(create_statement[quote_pos:], dialect) if not string_tokens or string_tokens[0].token_type != TokenType.STRING: return None sql = string_tokens[0].text else: sql = create_statement[start:] - try: - tokens = tokenize(sql) - except Exception: - return None + tokens = _tokenize_ddl_prefix(sql, dialect) depth = 0 for token in tokens: if token.token_type == TokenType.L_PAREN: @@ -425,3 +423,16 @@ def _extract_create_table_statement(create_statement: str) -> "str | None": if depth == 0: return sql[: token.end + 1] return None + + +def _tokenize_ddl_prefix(sql: str, dialect: "DialectType") -> "list[Token]": + """Retain complete tokens before an unsupported procedural wrapper suffix. + + Only a completed string literal or balanced table definition is consumed by + the caller. An unfinished token in the DDL cannot complete either boundary. + """ + tokenizer = Dialect.get_or_raise(dialect).tokenizer() + try: + return tokenizer.tokenize(sql) + except TokenError: + return tokenizer.tokens diff --git a/tests/unit/adapters/test_db2/test_config.py b/tests/unit/adapters/test_db2/test_config.py index 7cb53ca19..2b0957d58 100644 --- a/tests/unit/adapters/test_db2/test_config.py +++ b/tests/unit/adapters/test_db2/test_config.py @@ -31,6 +31,19 @@ FakeModules = tuple[FakeIbmDbModule, FakeIbmDbDbiModule] +@pytest.mark.parametrize("mode", ["Sync", "Async"]) +def test_connection_signature_namespace_preserves_native_types(mode: str) -> None: + native = pytest.importorskip("ibm_db_dbi") + config = Db2SyncConfig() if mode == "Sync" else Db2AsyncConfig() + connection_type = native.Connection if mode == "Sync" else native.AsyncConnection + cursor_type = native.Cursor if mode == "Sync" else native.AsyncCursor + namespace = config.get_signature_namespace() + + assert config.connection_type is connection_type + assert namespace[f"Db2{mode}Connection"] is connection_type + assert namespace["Db2RawCursor" if mode == "Sync" else "Db2AsyncRawCursor"] is cursor_type + + def test_parse_db2_dsn_cli_format() -> None: """Known CLI keywords map to canonical connection keys.""" dsn = "DATABASE=mytestdb;HOSTNAME=127.0.0.1;PORT=50000;PROTOCOL=TCPIP;UID=db2admin;PWD=secretpass;" diff --git a/tests/unit/adapters/test_duckdb/test_arrow_streaming.py b/tests/unit/adapters/test_duckdb/test_arrow_streaming.py index fd2923349..dd0bd74fe 100644 --- a/tests/unit/adapters/test_duckdb/test_arrow_streaming.py +++ b/tests/unit/adapters/test_duckdb/test_arrow_streaming.py @@ -16,8 +16,10 @@ @contextmanager -def _seed_driver() -> Iterator[DuckDBDriver]: - config = DuckDBConfig(connection_config={"database": ":memory:"}) +def _seed_driver(*, enable_uuid_conversion: bool = True) -> Iterator[DuckDBDriver]: + config = DuckDBConfig( + connection_config={"database": ":memory:"}, driver_features={"enable_uuid_conversion": enable_uuid_conversion} + ) with config.provide_session() as driver: driver.execute_script(""" CREATE OR REPLACE TABLE arrow_streaming (id INTEGER, name VARCHAR); @@ -98,18 +100,19 @@ def test_execute_many_bulk_path_survives_a_previous_driver() -> None: second_config.close_pool() -def test_select_arrow_path_restores_uuid_columns_only() -> None: +@pytest.mark.parametrize("enabled", [True, False]) +def test_select_arrow_path_restores_uuid_columns_only(enabled: bool) -> None: uuid_value = "550e8400-e29b-41d4-a716-446655440000" - with _seed_driver() as driver: + with _seed_driver(enable_uuid_conversion=enabled) as driver: driver.execute("CREATE OR REPLACE TABLE uuid_target (id UUID, text_id VARCHAR)") driver.execute("INSERT INTO uuid_target VALUES (?, ?)", uuid_value, uuid_value) - row = driver.select_one("SELECT id, text_id FROM uuid_target") - - assert isinstance(row["id"], UUID) - assert str(row["id"]) == uuid_value - assert row["text_id"] == uuid_value - assert isinstance(row["text_id"], str) + for _ in range(2): + row = driver.select_one("SELECT id, text_id FROM uuid_target WHERE text_id = ?", uuid_value) + assert isinstance(row["id"], UUID) + assert str(row["id"]) == uuid_value + assert row["text_id"] == uuid_value + assert isinstance(row["text_id"], str) def test_select_stream_native_only_reads_bounded_rows() -> None: @@ -128,9 +131,10 @@ def test_select_stream_native_only_reads_bounded_rows() -> None: ] -def test_select_stream_restores_uuid_columns_per_batch() -> None: +@pytest.mark.parametrize("enabled", [True, False]) +def test_select_stream_restores_uuid_columns_per_batch(enabled: bool) -> None: uuid_value = "550e8400-e29b-41d4-a716-446655440000" - with _seed_driver() as driver: + with _seed_driver(enable_uuid_conversion=enabled) as driver: driver.execute("CREATE OR REPLACE TABLE uuid_stream (id UUID, text_id VARCHAR)") driver.execute("INSERT INTO uuid_stream VALUES (?, ?), (?, ?)", uuid_value, uuid_value, uuid_value, uuid_value) diff --git a/tests/unit/adapters/test_mysqlconnector/test_config.py b/tests/unit/adapters/test_mysqlconnector/test_config.py index a693c566f..5eac4f9d1 100644 --- a/tests/unit/adapters/test_mysqlconnector/test_config.py +++ b/tests/unit/adapters/test_mysqlconnector/test_config.py @@ -25,6 +25,19 @@ from sqlspec.adapters.mysqlconnector.config import MysqlConnectorCursorParams, MysqlConnectorFailoverTarget +@pytest.mark.anyio +async def test_async_create_pool_uses_native_pool(monkeypatch: pytest.MonkeyPatch) -> None: + native_pool = pytest.importorskip("mysql.connector.aio.pooling").MySQLConnectionPool + initialize = AsyncMock() + monkeypatch.setattr(native_pool, "initialize_pool", initialize) + config = MysqlConnectorAsyncConfig(connection_config={"pool_name": "sqlspec", "pool_size": 2}) + + pool = await config.create_pool() + + assert isinstance(pool, native_pool) + initialize.assert_awaited_once_with() + + def test_sync_config_uses_connector_python_host_default_and_disables_local_infile() -> None: """SQLSpec should preserve the driver host default and close the local infile gate.""" config = MysqlConnectorSyncConfig() diff --git a/tests/unit/migrations/test_schema_ensure.py b/tests/unit/migrations/test_schema_ensure.py index fd2e65ac3..784ca0dd9 100644 --- a/tests/unit/migrations/test_schema_ensure.py +++ b/tests/unit/migrations/test_schema_ensure.py @@ -157,3 +157,40 @@ def test_schema_target_from_ddl_ignores_parentheses_in_strings_and_comments() -> assert [column.name for column in target.create_table.columns] == ["id", "delim", "label"] assert target.create_table.columns[2].not_null is True + + +@pytest.mark.parametrize( + ("dialect", "ddl", "columns"), + [ + ( + "tsql", + "IF OBJECT_ID(N'widgets', N'U') IS NULL BEGIN CREATE TABLE widgets (id INT, [O'Brien] VARCHAR(50)); END;", + ["id", "O'Brien"], + ), + ( + "oracle", + ( + "BEGIN EXECUTE IMMEDIATE 'CREATE TABLE widgets (id NUMBER, label VARCHAR2(50))'; " + "EXCEPTION WHEN OTHERS THEN raise_application_error(-20001, q'[can't create table]'); END;" + ), + ["id", "label"], + ), + ], +) +def test_schema_target_from_wrapped_ddl_preserves_dialect_quotes(dialect: str, ddl: str, columns: list[str]) -> None: + target = SchemaTarget.from_ddl("widgets", ddl, dialect=dialect) + + assert [column.name for column in target.create_table.columns] == columns + assert target.create_statement == ddl + + +@pytest.mark.parametrize( + "ddl", + [ + "BEGIN EXECUTE IMMEDIATE 'CREATE TABLE widgets (id NUMBER); END;", + "BEGIN EXECUTE IMMEDIATE 'CREATE TABLE widgets (id NUMBER, label VARCHAR2(50)'; END;", + ], +) +def test_schema_target_rejects_incomplete_wrapped_ddl(ddl: str) -> None: + with pytest.raises(ValueError, match="Unable to parse CREATE TABLE DDL"): + SchemaTarget.from_ddl("widgets", ddl, dialect="oracle")