From c3a240a83c5d27ec13d138c7b3ab3267b9d7622a Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Fri, 25 Sep 2026 00:17:24 +0000 Subject: [PATCH 01/10] refactor(adapters): clean up PostgreSQL and CockroachDB adapters - AsyncPG: remove dead execution helpers, use singleton RNG and exception chain depth limit, optimize connection pool configs and ADK session round-trips. - Psycopg: simplify connection and driver hierarchy, remove dead passthrough methods, add typed query conversion helper without byte encoding, streamline ADK and Litestar session stores. - PsqlPy: remove obsolete driver wrappers and query cache indirection, consolidate connection pool configs, optimize Arrow conversion helpers. - CockroachDB (AsyncPG & Psycopg): inherit from PostgreSQL base driver and exception handlers, fix follower read transaction commit leak, specialize data dictionaries with per-connection version caching, align multi-region configuration options. --- sqlspec/adapters/asyncpg/adk/store.py | 18 +- sqlspec/adapters/asyncpg/config.py | 43 ++- sqlspec/adapters/asyncpg/core.py | 53 ++- sqlspec/adapters/asyncpg/driver.py | 104 +++++- sqlspec/adapters/asyncpg/litestar/store.py | 27 +- sqlspec/adapters/cockroach_asyncpg/config.py | 10 +- sqlspec/adapters/cockroach_asyncpg/core.py | 38 ++- .../cockroach_asyncpg/data_dictionary.py | 50 ++- sqlspec/adapters/cockroach_asyncpg/driver.py | 70 ++-- .../adapters/cockroach_psycopg/__init__.py | 8 +- .../adapters/cockroach_psycopg/adk/store.py | 125 ++++--- sqlspec/adapters/cockroach_psycopg/config.py | 42 +-- sqlspec/adapters/cockroach_psycopg/core.py | 76 ++++- .../cockroach_psycopg/data_dictionary.py | 116 ++++--- sqlspec/adapters/cockroach_psycopg/driver.py | 138 +++----- .../cockroach_psycopg/litestar/store.py | 33 +- sqlspec/adapters/psqlpy/adk/store.py | 84 +++-- sqlspec/adapters/psqlpy/config.py | 25 +- sqlspec/adapters/psqlpy/core.py | 126 ++++++- sqlspec/adapters/psqlpy/driver.py | 312 +++++++++++++++++- sqlspec/adapters/psqlpy/litestar/store.py | 43 +-- sqlspec/adapters/psqlpy/type_converter.py | 32 +- sqlspec/adapters/psycopg/_typing.py | 4 + sqlspec/adapters/psycopg/adk/store.py | 121 ++++--- sqlspec/adapters/psycopg/config.py | 79 ++++- sqlspec/adapters/psycopg/core.py | 25 +- sqlspec/adapters/psycopg/driver.py | 104 +++--- sqlspec/adapters/psycopg/type_converter.py | 70 +++- 28 files changed, 1405 insertions(+), 571 deletions(-) diff --git a/sqlspec/adapters/asyncpg/adk/store.py b/sqlspec/adapters/asyncpg/adk/store.py index 8b8d84221..ac45083b2 100644 --- a/sqlspec/adapters/asyncpg/adk/store.py +++ b/sqlspec/adapters/asyncpg/adk/store.py @@ -106,20 +106,28 @@ async def create_session( INSERT INTO {self._session_table} (id, app_name, user_id, {self._owner_id_column_name}, state, create_time, update_time) VALUES ($1, $2, $3, $4, $5, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ - await conn.execute(sql, session_id, app_name, user_id, owner_id, state) + row = await conn.fetchrow(sql, session_id, app_name, user_id, owner_id, state) else: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time) VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ - await conn.execute(sql, session_id, app_name, user_id, state) + row = await conn.fetchrow(sql, session_id, app_name, user_id, state) - result = await self.get_session(app_name, user_id, session_id) - if result is None: + if row is None: msg = "Failed to fetch created session" raise RuntimeError(msg) - return result + return StoredSession( + id=row["id"], + app_name=row["app_name"], + user_id=row["user_id"], + state=row["state"], + create_time=row["create_time"], + update_time=row["update_time"], + ) async def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index f92543131..e04eeacbd 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -101,6 +101,7 @@ class AsyncpgConnectionConfig(TypedDict): connect_timeout: NotRequired[float] command_timeout: NotRequired[float] statement_cache_size: NotRequired[int] + pgbouncer: NotRequired[bool] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] server_settings: NotRequired["dict[str, str]"] @@ -191,6 +192,9 @@ class AsyncpgDriverFeatures(TypedDict): - "notify_queue": Durable queue plus a PostgreSQL notification wakeup hint - "poll_queue": Durable queue discovered by polling Defaults to "notify". + pgbouncer: Enable PgBouncer transaction-pooling compatibility mode. + Disables server-side prepared statement caching (statement_cache_size=0). + type_codecs: Optional list of custom type codec specifications to register. """ json_serializer: NotRequired["Callable[[Any], str]"] @@ -211,6 +215,8 @@ class AsyncpgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["notify", "notify_queue", "poll_queue"]] connection_instance: NotRequired["AsyncpgPool"] on_connection_create: NotRequired["Callable[[AsyncpgConnection], Awaitable[None]]"] + pgbouncer: NotRequired[bool] + type_codecs: NotRequired["list[dict[str, Any]]"] class _AsyncpgCloudSqlConnector: @@ -334,6 +340,7 @@ def __init__( self._user_connection_hook: Callable[[AsyncpgConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) + self._custom_type_codecs: list[dict[str, Any]] = list(features_dict.pop("type_codecs", None) or []) super().__init__( connection_config=build_connection_config(normalize_connection_config(connection_config)), @@ -451,12 +458,33 @@ def _setup_alloydb_connector(self, config: "dict[str, Any]") -> None: config["connect"] = _AsyncpgAlloydbConnector(self, user, password, database) + def register_type_codec( + self, + typename: str, + *, + schema: str = "public", + encoder: "Callable[..., Any] | None" = None, + decoder: "Callable[..., Any] | None" = None, + format: str = "text", + ) -> None: + """Register a custom type codec to be applied to all connections.""" + self._custom_type_codecs.append({ + "typename": typename, + "schema": schema, + "encoder": encoder, + "decoder": decoder, + "format": format, + }) + async def _create_pool(self) -> "Pool[Record]": """Create the actual async connection pool.""" config = { key: value for key, value in build_connection_config(self.connection_config).items() if value is not None } + if self.connection_config.get("pgbouncer") or self.driver_features.get("pgbouncer"): + config["statement_cache_size"] = 0 + if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): @@ -467,7 +495,7 @@ async def _create_pool(self) -> "Pool[Record]": return await asyncpg_create_pool(**config) async def _init_connection(self, connection: "AsyncpgConnection") -> None: - """Initialize connection with JSON codecs, pgvector support, and user callback. + """Initialize connection with JSON codecs, pgvector support, custom codecs, and user callback. Args: connection: AsyncPG connection to initialize. @@ -479,7 +507,6 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: decoder=self.driver_features.get("json_deserializer", from_json), ) - # Detect extensions on first connection, update dialect if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -500,7 +527,17 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: if self._pgvector_available: await register_pgvector_support(connection) - # Call user-provided callback after internal setup + for codec in self._custom_type_codecs: + codec_kwargs: dict[str, Any] = { + "schema": codec.get("schema", "public"), + "format": codec.get("format", "text"), + } + if codec.get("encoder") is not None: + codec_kwargs["encoder"] = codec["encoder"] + if codec.get("decoder") is not None: + codec_kwargs["decoder"] = codec["decoder"] + await connection.set_type_codec(codec["typename"], **codec_kwargs) + if self._user_connection_hook is not None: await self._user_connection_hook(connection) diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index ef7620600..62a494216 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -4,7 +4,7 @@ import datetime import re from collections.abc import Sized -from typing import TYPE_CHECKING, Any, Final, NamedTuple +from typing import TYPE_CHECKING, Any, Final, NamedTuple, cast from sqlspec.adapters.asyncpg._typing import asyncpg_module as asyncpg from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile @@ -297,10 +297,17 @@ def parse_status(status: Any) -> int: if not status or not isinstance(status, str): return 0 - match = ASYNC_PG_STATUS_REGEX.match(status.strip()) + stripped = status.strip() + last_space = stripped.rfind(" ") + if last_space != -1: + token = stripped[last_space + 1 :] + if token.isdigit(): + return int(token) + + match = ASYNC_PG_STATUS_REGEX.match(stripped) if match: groups = match.groups() - if len(groups) >= EXPECTED_REGEX_GROUPS: + if len(groups) >= EXPECTED_REGEX_GROUPS and groups[-1]: try: return int(groups[-1]) except (ValueError, IndexError): @@ -452,12 +459,13 @@ async def _start(self) -> None: self._transaction = None raise - async def fetch_chunk(self) -> "list[dict[str, Any]]": + async def fetch_chunk(self) -> "list[Any]": handler = self._driver.handle_database_exceptions() records = await self._driver._run_with_exception_handler(handler, self._cursor.fetch, self._chunk_size) self._driver._check_pending_exception(handler) - assert records is not None - return [dict(record) for record in records] + if records is None: + return [] + return cast("list[Any]", records) async def close(self, error: bool = False) -> None: self._cursor = None @@ -529,12 +537,20 @@ def _encode_json_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: return str(encoded).encode("utf-8") -def _decode_json_payload(value: Any, decoder: "Callable[[str], Any]") -> Any: +def _decode_json_payload(value: Any, decoder: "Callable[..., Any]") -> Any: + """Decode JSON binary or string payload with zero-copy decoding when possible.""" if isinstance(value, str): return decoder(value) if isinstance(value, memoryview): - value = value.tobytes() - return decoder(bytes(value).decode("utf-8")) + raw_bytes = value.tobytes() + elif isinstance(value, (bytes, bytearray)): + raw_bytes = bytes(value) + else: + raw_bytes = bytes(value) + try: + return decoder(raw_bytes) + except (TypeError, UnicodeDecodeError): + return decoder(raw_bytes.decode("utf-8")) def _encode_jsonb_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: @@ -544,15 +560,22 @@ def _encode_jsonb_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: return _JSONB_BINARY_VERSION + payload -def _decode_jsonb_payload(value: Any, decoder: "Callable[[str], Any]") -> Any: +def _decode_jsonb_payload(value: Any, decoder: "Callable[..., Any]") -> Any: + """Decode JSONB binary or string payload stripping version prefix when present.""" if isinstance(value, str): return decoder(value) if isinstance(value, memoryview): - value = value.tobytes() - payload = bytes(value) - if payload.startswith(_JSONB_BINARY_VERSION): - payload = payload[1:] - return decoder(payload.decode("utf-8")) + raw_bytes = value.tobytes() + elif isinstance(value, (bytes, bytearray)): + raw_bytes = bytes(value) + else: + raw_bytes = bytes(value) + if raw_bytes.startswith(_JSONB_BINARY_VERSION): + raw_bytes = raw_bytes[1:] + try: + return decoder(raw_bytes) + except (TypeError, UnicodeDecodeError): + return decoder(raw_bytes.decode("utf-8")) def _create_postgres_error( diff --git a/sqlspec/adapters/asyncpg/driver.py b/sqlspec/adapters/asyncpg/driver.py index eccab9969..eee6b7055 100644 --- a/sqlspec/adapters/asyncpg/driver.py +++ b/sqlspec/adapters/asyncpg/driver.py @@ -7,6 +7,7 @@ from io import BytesIO from typing import TYPE_CHECKING, Any, Final, cast +from mypy_extensions import mypyc_attr from sqlglot import exp, parse_one from sqlglot.errors import ParseError @@ -87,6 +88,7 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class AsyncpgDriver(AsyncDriverAdapterBase): """AsyncPG PostgreSQL driver for async database operations. @@ -128,9 +130,24 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") """ sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) params: tuple[Any, ...] = cast("tuple[Any, ...]", prepared_parameters) if prepared_parameters else () + execution_args = statement.statement_config.execution_args or {} + driver_args = self.statement_config.execution_args or {} + command_timeout = ( + execution_args.get("timeout") + or execution_args.get("command_timeout") + or driver_args.get("timeout") + or driver_args.get("command_timeout") + ) if statement.returns_rows(): - records = await cursor.fetch(sql, *params) if params else await cursor.fetch(sql) + if command_timeout is not None: + records = ( + await cursor.fetch(sql, *params, timeout=command_timeout) + if params + else await cursor.fetch(sql, timeout=command_timeout) + ) + else: + records = await cursor.fetch(sql, *params) if params else await cursor.fetch(sql) data, column_names = collect_rows(records) return self.create_execution_result( @@ -142,7 +159,17 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") row_format="record", ) - result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) + if command_timeout is not None: + result = ( + await cursor.execute(sql, *params, timeout=command_timeout) + if params + else await cursor.execute(sql, timeout=command_timeout) + ) + else: + result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) + + if statement.operation_type in {"CREATE", "ALTER", "DROP", "TRUNCATE"}: + self.invalidate_prepared_statements() affected_rows = parse_status(result) @@ -190,6 +217,8 @@ async def dispatch_execute_script(self, cursor: "AsyncpgConnection", statement: last_result = result successful_count += 1 + self.invalidate_prepared_statements() + return self.create_execution_result( last_result, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) @@ -230,8 +259,7 @@ async def begin(self) -> None: await transaction.start() except AsyncpgPostgresError as e: self._release_failed_transaction_claim(transaction) - msg = f"Failed to begin async transaction: {e}" - raise SQLSpecError(msg) from e + raise create_mapped_exception(e) from e self._transaction = transaction def _release_failed_transaction_claim(self, transaction: Any) -> None: @@ -250,8 +278,7 @@ async def commit(self) -> None: else: await self.connection.execute("COMMIT") except AsyncpgPostgresError as e: - msg = f"Failed to commit async transaction: {e}" - raise SQLSpecError(msg) from e + raise create_mapped_exception(e) from e async def rollback(self) -> None: """Rollback the current transaction.""" @@ -263,17 +290,18 @@ async def rollback(self) -> None: else: await self.connection.execute("ROLLBACK") except AsyncpgPostgresError as e: - msg = f"Failed to rollback async transaction: {e}" - raise SQLSpecError(msg) from e + raise create_mapped_exception(e) from e async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" + self.invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET LOCAL search_path TO {quoted_schema}, "$user", public') async def set_migration_non_transactional_schema(self, schema: str) -> None: """Set the PostgreSQL search path for non-transactional migration SQL.""" + self.invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET search_path TO {quoted_schema}, "$user", public') @@ -364,11 +392,16 @@ async def load_from_arrow( except AsyncpgPostgresError as exc: msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - columns, records = self._arrow_table_to_rows(arrow_table) - if records: - await self.connection.copy_records_to_table( - table_name, records=records, columns=columns, schema_name=schema_name - ) + + import pyarrow as pa + + for batch in arrow_table.to_batches(): + batch_table = pa.Table.from_batches([batch]) + columns, records = self._arrow_table_to_rows(batch_table) + if records: + await self.connection.copy_records_to_table( + table_name, records=records, columns=columns, schema_name=schema_name + ) telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table self._attach_partition_telemetry(telemetry_payload, partitioner) @@ -381,8 +414,9 @@ async def load_from_records( *, columns: "list[str] | None" = None, overwrite: bool = False, + batch_size: int = 1000, ) -> "StorageBridgeJob": - """Load mapping or positional records directly with binary COPY.""" + """Load mapping or positional records directly with binary COPY in batches.""" self._require_capability("arrow_import_enabled") materialized = list(records) if not materialized: @@ -425,9 +459,13 @@ async def load_from_records( except AsyncpgPostgresError as exc: msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - await self.connection.copy_records_to_table( - table_name, records=copy_rows, columns=resolved_columns, schema_name=schema_name - ) + + for i in range(0, len(copy_rows), batch_size): + chunk = copy_rows[i : i + batch_size] + await self.connection.copy_records_to_table( + table_name, records=chunk, columns=resolved_columns, schema_name=schema_name + ) + telemetry_payload: StorageTelemetry = { "bytes_processed": 0, "destination": table, @@ -436,6 +474,30 @@ async def load_from_records( } return self._storage_job(telemetry_payload) + async def copy_from_table( + self, + table: str, + output: Any, + *, + columns: "list[str] | None" = None, + schema_name: "str | None" = None, + format: str = "text", + delimiter: str = "\t", + null: str = "\\N", + ) -> None: + """Export table contents to output stream or file using PostgreSQL COPY TO STDOUT.""" + table_name, resolved_schema, _ = self._copy_target(table) + schema = schema_name or resolved_schema + await self.connection.copy_from_table( + table_name, + output=output, + columns=columns, + schema_name=schema, + format=format, + delimiter=delimiter, + null=null, + ) + async def load_from_storage( self, table: str, @@ -592,7 +654,15 @@ def _connection_in_transaction(self) -> bool: """Check if connection is in transaction.""" return bool(self.connection.is_in_transaction()) + def invalidate_prepared_statements(self) -> None: + """Clear cached prepared statements.""" + self._prepared_statements.clear() + async def _get_prepared_statement(self, sql: str) -> "AsyncpgPreparedStatement": + """Get or prepare a statement with LRU caching.""" + if self.driver_features.get("pgbouncer"): + return cast("AsyncpgPreparedStatement", await self.connection.prepare(sql)) + cached = self._prepared_statements.get(sql) if cached is not None: self._prepared_statements.move_to_end(sql) diff --git a/sqlspec/adapters/asyncpg/litestar/store.py b/sqlspec/adapters/asyncpg/litestar/store.py index 238aa0ebb..e60e1c98f 100644 --- a/sqlspec/adapters/asyncpg/litestar/store.py +++ b/sqlspec/adapters/asyncpg/litestar/store.py @@ -81,6 +81,23 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by Returns: Session data as bytes if found and not expired, None otherwise. """ + if renew_for is not None: + new_expires_at = self._calculate_expires_at(renew_for) + if new_expires_at is not None: + update_sql = f""" + UPDATE {self._table_name} + SET expires_at = CASE WHEN expires_at IS NOT NULL THEN $1 ELSE expires_at END, + updated_at = CURRENT_TIMESTAMP + WHERE session_id = $2 + AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) + RETURNING data + """ + async with self._config.provide_connection() as conn: + row = await conn.fetchrow(update_sql, new_expires_at, key) + if row is None: + return None + return bytes(row["data"]) + sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = $1 @@ -93,16 +110,6 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by if row is None: return None - if renew_for is not None and row["expires_at"] is not None: - new_expires_at = self._calculate_expires_at(renew_for) - if new_expires_at is not None: - update_sql = f""" - UPDATE {self._table_name} - SET expires_at = $1, updated_at = CURRENT_TIMESTAMP - WHERE session_id = $2 - """ - await conn.execute(update_sql, new_expires_at, key) - return bytes(row["data"]) async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: diff --git a/sqlspec/adapters/cockroach_asyncpg/config.py b/sqlspec/adapters/cockroach_asyncpg/config.py index 158eb7942..72fcaf9d4 100644 --- a/sqlspec/adapters/cockroach_asyncpg/config.py +++ b/sqlspec/adapters/cockroach_asyncpg/config.py @@ -7,7 +7,6 @@ from sqlspec.adapters.asyncpg.core import ( apply_driver_features, - build_connection_config, default_statement_config, register_json_codecs, register_pgvector_support, @@ -20,7 +19,7 @@ from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgRecord as Record from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_connect as asyncpg_connect from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_create_pool as asyncpg_create_pool -from sqlspec.adapters.cockroach_asyncpg.core import validate_follower_read_staleness +from sqlspec.adapters.cockroach_asyncpg.core import build_connection_config, validate_follower_read_staleness from sqlspec.adapters.cockroach_asyncpg.driver import CockroachAsyncpgDriver, CockroachAsyncpgExceptionHandler from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs from sqlspec.core.capabilities import TypeCoercionCapabilities @@ -80,6 +79,10 @@ class CockroachAsyncpgConnectionConfig(TypedDict): timeout: NotRequired[float] connect_timeout: NotRequired[float] command_timeout: NotRequired[float] + application_name: NotRequired[str] + gateway_region: NotRequired[str] + default_transaction_use_follower_reads: NotRequired[bool] + results_buffer_size: NotRequired[int] statement_cache_size: NotRequired[int] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] @@ -175,7 +178,6 @@ class CockroachAsyncpgDriverFeatures(TypedDict): class _CockroachAsyncpgSessionFactory(AsyncPoolSessionFactory): """Uses pool.acquire() context manager pattern instead of direct acquire/release.""" - # _connection inherited from AsyncPoolSessionFactory.__slots__ is never written; this class uses _ctx exclusively via the pool.acquire() context manager pattern. __slots__ = ("_ctx",) def __init__(self, config: "CockroachAsyncpgConfig") -> None: @@ -209,7 +211,7 @@ class CockroachAsyncpgConfig( """Configuration for CockroachDB using AsyncPG.""" driver_type: "ClassVar[type[CockroachAsyncpgDriver]]" = CockroachAsyncpgDriver - connection_type: "ClassVar[type[CockroachAsyncpgConnection]]" = CockroachAsyncpgConnection # type: ignore[assignment] + connection_type: "ClassVar[type[CockroachAsyncpgConnection]]" = cast("Any", CockroachAsyncpgConnection) supports_transactional_ddl: "ClassVar[bool]" = False supports_migration_schemas: "ClassVar[bool]" = True supports_native_arrow_export: "ClassVar[bool]" = True diff --git a/sqlspec/adapters/cockroach_asyncpg/core.py b/sqlspec/adapters/cockroach_asyncpg/core.py index fefa80e03..b5ab58983 100644 --- a/sqlspec/adapters/cockroach_asyncpg/core.py +++ b/sqlspec/adapters/cockroach_asyncpg/core.py @@ -1,7 +1,7 @@ """CockroachDB AsyncPG adapter helpers.""" -import random import re +import secrets from typing import TYPE_CHECKING, Any, Final, cast from mypy_extensions import mypyc_attr @@ -19,6 +19,7 @@ __all__ = ( "CockroachAsyncpgRetryConfig", + "build_connection_config", "build_native_export", "build_native_import", "calculate_backoff_seconds", @@ -29,11 +30,12 @@ "validate_follower_read_staleness", ) -# Retry configuration defaults (module-level for mypyc compatibility) _DEFAULT_MAX_RETRIES: Final[int] = 10 _DEFAULT_BASE_DELAY_MS: Final[float] = 50.0 _DEFAULT_MAX_DELAY_MS: Final[float] = 5000.0 _DEFAULT_ENABLE_LOGGING: Final[bool] = True +_MAX_EXCEPTION_CHAIN_DEPTH: Final[int] = 16 +_RNG: Final = secrets.SystemRandom() @mypyc_attr(allow_interpreted_subclasses=False) @@ -65,6 +67,28 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachAsyncp ) +def build_connection_config(config: "dict[str, Any]") -> "dict[str, Any]": + """Prepare CockroachDB AsyncPG connection config, extracting multi-region server settings.""" + from sqlspec.adapters.asyncpg.core import build_connection_config as asyncpg_build_connection_config + + result = asyncpg_build_connection_config(config) + server_settings = dict(result.get("server_settings") or {}) + if "application_name" in result: + server_settings.setdefault("application_name", str(result.pop("application_name"))) + if "gateway_region" in result: + server_settings.setdefault("gateway_region", str(result.pop("gateway_region"))) + if "default_transaction_use_follower_reads" in result: + val = result.pop("default_transaction_use_follower_reads") + server_settings.setdefault( + "default_transaction_use_follower_reads", "on" if val is True or str(val).lower() == "on" else "off" + ) + if "results_buffer_size" in result: + server_settings.setdefault("results_buffer_size", str(result.pop("results_buffer_size"))) + if server_settings: + result["server_settings"] = server_settings + return result + + def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. @@ -85,17 +109,17 @@ def is_retryable_error(error: BaseException) -> bool: Returns: True when the transaction should be retried. """ - seen: set[int] = set() + depth = 0 current: BaseException | None = error - while current is not None and id(current) not in seen: - seen.add(id(current)) + while current is not None and depth < _MAX_EXCEPTION_CHAIN_DEPTH: if isinstance(current, SerializationConflictError): return True if has_sqlstate(current) and str(current.sqlstate) == "40001": return True if not isinstance(current, SQLSpecError): return False - current = cast("BaseException | None", cast("Any", current).__cause__) + current = cast("BaseException | None", getattr(current, "__cause__", None)) + depth += 1 return False @@ -109,7 +133,7 @@ def calculate_backoff_seconds(attempt: int, config: "CockroachAsyncpgRetryConfig capped_ms: float = min(config.base_delay_ms * (2**attempt), config.max_delay_ms) if capped_ms <= 0.0: return 0.0 - return random.uniform(capped_ms / 2.0, capped_ms) / 1000.0 # noqa: S311 + return _RNG.uniform(capped_ms / 2.0, capped_ms) / 1000.0 _STALENESS_LITERAL: Final[re.Pattern[str]] = re.compile(r"'[^'\\;]+'") diff --git a/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py b/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py index 6839999c3..459028332 100644 --- a/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py +++ b/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py @@ -3,6 +3,7 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar +from sqlspec.adapters.asyncpg.data_dictionary import AsyncpgDataDictionary from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -27,10 +28,9 @@ unsupported_system_metadata_capability, ) from sqlspec.data_dictionary.dialects.cockroachdb import resolve_cockroachdb_json_type -from sqlspec.driver import AsyncDataDictionaryBase if TYPE_CHECKING: - from sqlspec.adapters.cockroach_asyncpg.driver import CockroachAsyncpgDriver + from sqlspec.adapters.asyncpg.driver import AsyncpgDriver from sqlspec.core import SQL __all__ = ("CockroachAsyncpgDataDictionary",) @@ -55,19 +55,19 @@ _COCKROACH_SUPPORTED_DOMAINS = frozenset(_COCKROACH_METADATA_DOMAINS) - {"crdb_internal", "system"} -class CockroachAsyncpgDataDictionary(AsyncDataDictionaryBase): +class CockroachAsyncpgDataDictionary(AsyncpgDataDictionary): """CockroachDB async data dictionary (AsyncPG).""" dialect: ClassVar[str] = "cockroachdb" async def get_metadata_capabilities( - self, driver: "CockroachAsyncpgDriver", domains: Sequence[str] | None = None + self, driver: "AsyncpgDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) async def get_system_metadata_capabilities( - self, driver: "CockroachAsyncpgDriver", domains: Sequence[str] | None = None + self, driver: "AsyncpgDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -75,7 +75,7 @@ async def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) async def _select_domain( - self, driver: "CockroachAsyncpgDriver", domain: str, query_name: str, **parameters: Any + self, driver: "AsyncpgDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -83,17 +83,15 @@ async def _select_domain( rows = await driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - async def get_schemas(self, driver: "CockroachAsyncpgDriver") -> MetadataResult: + async def get_schemas(self, driver: "AsyncpgDriver") -> MetadataResult: """Get schema metadata.""" return await self._select_domain(driver, "schemas", "list") - async def get_objects(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_objects(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return await self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_table_details( - self, driver: "CockroachAsyncpgDriver", table: str, schema: str | None = None - ) -> MetadataResult: + async def get_table_details(self, driver: "AsyncpgDriver", table: str, schema: str | None = None) -> MetadataResult: """Get rich table metadata.""" return await self._select_domain( driver, @@ -104,7 +102,7 @@ async def get_table_details( ) async def get_constraints( - self, driver: "CockroachAsyncpgDriver", table: str | None = None, schema: str | None = None + self, driver: "AsyncpgDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -112,16 +110,16 @@ async def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - async def get_views(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_views(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return await self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_routines(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_routines(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) async def get_privileges( - self, driver: "CockroachAsyncpgDriver", object_name: str | None = None, schema: str | None = None + self, driver: "AsyncpgDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -130,7 +128,7 @@ async def get_privileges( ) async def get_dependencies( - self, driver: "CockroachAsyncpgDriver", object_name: str | None = None, schema: str | None = None + self, driver: "AsyncpgDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -140,7 +138,7 @@ async def get_dependencies( async def get_ddl( self, - driver: "CockroachAsyncpgDriver", + driver: "AsyncpgDriver", object_name: str, schema: str | None = None, *, @@ -166,7 +164,7 @@ async def get_ddl( ) async def get_system_metadata( - self, driver: "CockroachAsyncpgDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "AsyncpgDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -174,8 +172,8 @@ async def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - async def get_version(self, driver: "CockroachAsyncpgDriver") -> "VersionInfo | None": - driver_id = id(driver) + async def get_version(self, driver: "AsyncpgDriver") -> "VersionInfo | None": + driver_id = id(driver.connection) if hasattr(driver, "connection") else id(driver) if driver_id in self._version_fetch_attempted: return self._version_cache.get(driver_id) @@ -195,17 +193,17 @@ async def get_version(self, driver: "CockroachAsyncpgDriver") -> "VersionInfo | self.cache_version(driver_id, version_info) return version_info - async def get_feature_flag(self, driver: "CockroachAsyncpgDriver", feature: str) -> bool: + async def get_feature_flag(self, driver: "AsyncpgDriver", feature: str) -> bool: version_info = await self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - async def get_optimal_type(self, driver: "CockroachAsyncpgDriver", type_category: str) -> str: + async def get_optimal_type(self, driver: "AsyncpgDriver", type_category: str) -> str: config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(await self.get_version(driver)) return config.get_optimal_type(type_category) - async def get_tables(self, driver: "CockroachAsyncpgDriver", schema: "str | None" = None) -> "list[TableMetadata]": + async def get_tables(self, driver: "AsyncpgDriver", schema: "str | None" = None) -> "list[TableMetadata]": schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") return await driver.select( @@ -216,7 +214,7 @@ async def get_tables(self, driver: "CockroachAsyncpgDriver", schema: "str | None ) async def get_columns( - self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": schema_name = self.resolve_schema(schema) if table is None: @@ -239,7 +237,7 @@ async def get_columns( ) async def get_indexes( - self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": schema_name = self.resolve_schema(schema) if table is None: @@ -261,7 +259,7 @@ async def get_indexes( ) async def get_foreign_keys( - self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": schema_name = self.resolve_schema(schema) if table is None: diff --git a/sqlspec/adapters/cockroach_asyncpg/driver.py b/sqlspec/adapters/cockroach_asyncpg/driver.py index 578b8f72e..82379cc95 100644 --- a/sqlspec/adapters/cockroach_asyncpg/driver.py +++ b/sqlspec/adapters/cockroach_asyncpg/driver.py @@ -1,10 +1,11 @@ """CockroachDB AsyncPG driver implementation.""" import asyncio +import contextlib from typing import TYPE_CHECKING, Any, TypeVar, cast from sqlspec.adapters.asyncpg.core import create_mapped_exception, driver_profile -from sqlspec.adapters.asyncpg.driver import AsyncpgDriver +from sqlspec.adapters.asyncpg.driver import AsyncpgDriver, AsyncpgExceptionHandler from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgPostgresError, CockroachAsyncpgSessionContext from sqlspec.adapters.cockroach_asyncpg.core import ( CockroachAsyncpgRetryConfig, @@ -19,7 +20,6 @@ ) from sqlspec.adapters.cockroach_asyncpg.data_dictionary import CockroachAsyncpgDataDictionary from sqlspec.core import SQL, register_driver_profile -from sqlspec.driver import BaseAsyncExceptionHandler from sqlspec.utils.logging import get_logger from sqlspec.utils.type_guards import has_sqlstate @@ -37,7 +37,7 @@ _T = TypeVar("_T") -class CockroachAsyncpgExceptionHandler(BaseAsyncExceptionHandler): +class CockroachAsyncpgExceptionHandler(AsyncpgExceptionHandler): """Async context manager for CockroachDB AsyncPG exceptions.""" __slots__ = () @@ -66,7 +66,6 @@ def __init__( self._retry_config = CockroachAsyncpgRetryConfig.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 select_to_storage( @@ -216,61 +215,34 @@ async def run_transaction_with_retry(self, operation: "Callable[[], Awaitable[_T attempt += 1 async def dispatch_execute(self, cursor: Any, statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_impl(cursor, statement) - - async def dispatch_execute_many(self, cursor: Any, statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_many_impl(cursor, statement) - - async def dispatch_execute_script(self, cursor: Any, statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_script_impl(cursor, statement) - - def handle_database_exceptions(self) -> "CockroachAsyncpgExceptionHandler": # type: ignore[override] + opened_txn = False + if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): + await self.begin() + opened_txn = True + try: + return await super().dispatch_execute(cursor, statement) + finally: + if opened_txn: + with contextlib.suppress(Exception): + await self.commit() + + def handle_database_exceptions(self) -> "CockroachAsyncpgExceptionHandler": return CockroachAsyncpgExceptionHandler() @property - def data_dictionary(self) -> "CockroachAsyncpgDataDictionary": # type: ignore[override] + def data_dictionary(self) -> "CockroachAsyncpgDataDictionary": if self._data_dictionary is None: - # Intentionally assign CockroachDB-specific data dictionary to parent slot - object.__setattr__(self, "_data_dictionary", CockroachAsyncpgDataDictionary()) + self._data_dictionary = CockroachAsyncpgDataDictionary() return cast("CockroachAsyncpgDataDictionary", self._data_dictionary) + def _follower_reads_enabled(self) -> bool: + return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) + async def _apply_follower_reads(self) -> None: - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: + if not self._follower_reads_enabled() or not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) await self.connection.execute(f"SET TRANSACTION AS OF SYSTEM TIME {staleness}") - async def _begin_follower_read_transaction(self) -> None: - """Open the transaction a follower read needs so the staleness clause can lead it. - - A statement run outside a transaction gets its own implicit one, which - the clause could not precede, so a read opens a transaction here when the - caller has not already done so. - """ - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: - return - if self._connection_in_transaction(): - return - await self.begin() - - async def _dispatch_execute_impl(self, cursor: "CockroachAsyncpgConnection", statement: SQL) -> "ExecutionResult": - if statement.returns_rows(): - await self._begin_follower_read_transaction() - return await super().dispatch_execute(cursor, statement) - - async def _dispatch_execute_many_impl( - self, cursor: "CockroachAsyncpgConnection", statement: SQL - ) -> "ExecutionResult": - return await AsyncpgDriver.dispatch_execute_many(self, cursor, statement) - - async def _dispatch_execute_script_impl( - self, cursor: "CockroachAsyncpgConnection", statement: SQL - ) -> "ExecutionResult": - return await AsyncpgDriver.dispatch_execute_script(self, cursor, statement) - register_driver_profile("cockroach_asyncpg", driver_profile) diff --git a/sqlspec/adapters/cockroach_psycopg/__init__.py b/sqlspec/adapters/cockroach_psycopg/__init__.py index 963f80a06..a36a39fd9 100644 --- a/sqlspec/adapters/cockroach_psycopg/__init__.py +++ b/sqlspec/adapters/cockroach_psycopg/__init__.py @@ -12,7 +12,12 @@ CockroachPsycopgSyncConfig, build_connection_config, ) -from sqlspec.adapters.cockroach_psycopg.core import CockroachPsycopgRetryConfig, build_statement_config, driver_profile +from sqlspec.adapters.cockroach_psycopg.core import ( + CockroachPsycopgRetryConfig, + as_query, + build_statement_config, + driver_profile, +) from sqlspec.adapters.cockroach_psycopg.driver import ( CockroachPsycopgAsyncDriver, CockroachPsycopgAsyncExceptionHandler, @@ -35,6 +40,7 @@ "CockroachPsycopgSyncExceptionHandler", "CockroachPsycopgSyncSessionContext", "CockroachSyncConnection", + "as_query", "build_connection_config", "build_statement_config", "driver_profile", diff --git a/sqlspec/adapters/cockroach_psycopg/adk/store.py b/sqlspec/adapters/cockroach_psycopg/adk/store.py index fc868818e..58e58f480 100644 --- a/sqlspec/adapters/cockroach_psycopg/adk/store.py +++ b/sqlspec/adapters/cockroach_psycopg/adk/store.py @@ -8,6 +8,7 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_dict_row as dict_row from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_errors as errors from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_sql as pg_sql +from sqlspec.adapters.cockroach_psycopg.core import as_query from sqlspec.config import ADKConfig from sqlspec.extensions.adk import ( BaseAsyncADKStore, @@ -229,24 +230,34 @@ async def create_session( sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, {self._owner_id_column_name}, state, create_time, update_time) VALUES (%s, %s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ params = (session_id, app_name, user_id, owner_id, state_json) else: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time) VALUES (%s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ params = (session_id, app_name, user_id, state_json) async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) + row = await cur.fetchone() await conn.commit() - result = await self.get_session(app_name, user_id, session_id) - if result is None: + if row is None: msg = "Session creation failed" raise RuntimeError(msg) - return result + + return StoredSession( + id=row["id"], + app_name=row["app_name"], + user_id=row["user_id"], + state=row["state"], + create_time=row["create_time"], + update_time=row["update_time"], + ) async def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None @@ -267,7 +278,7 @@ async def get_session( try: 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)) + await cur.execute(as_query(sql), (app_name, user_id, session_id)) row = await cur.fetchone() if row is None: @@ -292,7 +303,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(sql.encode(), (Jsonb(state), app_name, user_id, session_id)) + await cur.execute(as_query(sql), (Jsonb(state), app_name, user_id, session_id)) await conn.commit() async def list_sessions( @@ -315,7 +326,7 @@ async def list_sessions( try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) rows = await cur.fetchall() except errors.UndefinedTable: return [] @@ -336,7 +347,7 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> sql = f"DELETE FROM {self._session_table} WHERE app_name = %s AND user_id = %s AND id = %s" 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)) + await cur.execute(as_query(sql), (app_name, user_id, session_id)) await conn.commit() async def append_event(self, event_record: StoredEvent) -> None: @@ -350,7 +361,7 @@ async def append_event(self, event_record: StoredEvent) -> None: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute( - sql.encode(), + as_query(sql), ( event_record["id"], event_record["app_name"], @@ -399,7 +410,7 @@ async def append_event_and_update_state( async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: await cur.execute( - insert_sql.encode(), + as_query(insert_sql), ( event_record["id"], event_record["app_name"], @@ -410,14 +421,14 @@ async def append_event_and_update_state( jsonb_value, ), ) - await cur.execute(update_sql.encode(), (Jsonb(state), app_name, user_id, session_id)) + await cur.execute(as_query(update_sql), (Jsonb(state), app_name, user_id, session_id)) row = await cur.fetchone() if row is None: _raise_missing_session(session_id) if app_state is not None: - await cur.execute(app_upsert_sql.encode(), (app_name, Jsonb(app_state))) + await cur.execute(as_query(app_upsert_sql), (app_name, Jsonb(app_state))) if user_state is not None: - await cur.execute(user_upsert_sql.encode(), (app_name, user_id, Jsonb(user_state))) + await cur.execute(as_query(user_upsert_sql), (app_name, user_id, Jsonb(user_state))) except Exception: await conn.rollback() raise @@ -465,7 +476,7 @@ async def get_events( try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), tuple(params)) + await cur.execute(as_query(sql), tuple(params)) rows = await cur.fetchall() except errors.UndefinedTable: return [] @@ -493,7 +504,7 @@ async def delete_expired_events(self, before: "datetime", app_name: "str | None" try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -509,7 +520,7 @@ async def delete_idle_sessions(self, updated_before: "datetime", app_name: "str try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -525,7 +536,7 @@ async def delete_idle_user_states(self, updated_before: "datetime", app_name: "s try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -536,7 +547,7 @@ async def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), (app_name,)) + await cur.execute(as_query(sql), (app_name,)) row = await cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -547,7 +558,7 @@ async def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | try: 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)) + await cur.execute(as_query(sql), (app_name, user_id)) row = await cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -560,7 +571,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(sql.encode(), (app_name, Jsonb(state))) + await cur.execute(as_query(sql), (app_name, Jsonb(state))) await conn.commit() async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: @@ -570,7 +581,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(sql.encode(), (app_name, user_id, Jsonb(state))) + await cur.execute(as_query(sql), (app_name, user_id, Jsonb(state))) await conn.commit() async def get_metadata(self, key: str) -> "str | None": @@ -578,7 +589,7 @@ async def get_metadata(self, key: str) -> "str | None": try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), (key,)) + await cur.execute(as_query(sql), (key,)) row = await cur.fetchone() return row["value"] if row is not None else None except errors.UndefinedTable: @@ -591,7 +602,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(sql.encode(), (key, value)) + await cur.execute(as_query(sql), (key, value)) await conn.commit() async def _sessions_table_ddl(self) -> str: @@ -705,24 +716,34 @@ def create_session( sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, {self._owner_id_column_name}, state, create_time, update_time) VALUES (%s, %s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ params = (session_id, app_name, user_id, owner_id, state_json) else: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time) VALUES (%s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ params = (session_id, app_name, user_id, state_json) with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) + row = cur.fetchone() conn.commit() - result = self.get_session(app_name, user_id, session_id) - if result is None: + if row is None: msg = "Session creation failed" raise RuntimeError(msg) - return result + + return StoredSession( + id=row["id"], + app_name=row["app_name"], + user_id=row["user_id"], + state=row["state"], + create_time=row["create_time"], + update_time=row["update_time"], + ) def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None @@ -744,7 +765,7 @@ def get_session( try: 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)) + cur.execute(as_query(sql), (app_name, user_id, session_id)) row = cur.fetchone() if row is None: @@ -770,7 +791,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(sql.encode(), (Jsonb(state), app_name, user_id, session_id)) + cur.execute(as_query(sql), (Jsonb(state), app_name, user_id, session_id)) conn.commit() def list_sessions( @@ -794,7 +815,7 @@ def list_sessions( try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) rows = cur.fetchall() return [ @@ -816,7 +837,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 = %s AND user_id = %s AND id = %s" 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)) + cur.execute(as_query(sql), (app_name, user_id, session_id)) conn.commit() def append_event(self, event_record: StoredEvent) -> None: @@ -861,7 +882,7 @@ def append_event_and_update_state( with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: cur.execute( - insert_sql.encode(), + as_query(insert_sql), ( event_record["id"], event_record["app_name"], @@ -872,14 +893,14 @@ def append_event_and_update_state( jsonb_value, ), ) - cur.execute(update_sql.encode(), (Jsonb(state), app_name, user_id, session_id)) + cur.execute(as_query(update_sql), (Jsonb(state), app_name, user_id, session_id)) row = cur.fetchone() if row is None: _raise_missing_session(session_id) if app_state is not None: - cur.execute(app_upsert_sql.encode(), (app_name, Jsonb(app_state))) + cur.execute(as_query(app_upsert_sql), (app_name, Jsonb(app_state))) if user_state is not None: - cur.execute(user_upsert_sql.encode(), (app_name, user_id, Jsonb(user_state))) + cur.execute(as_query(user_upsert_sql), (app_name, user_id, Jsonb(user_state))) except Exception: conn.rollback() raise @@ -928,7 +949,7 @@ def get_events( try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), tuple(params)) + cur.execute(as_query(sql), tuple(params)) rows = cur.fetchall() return [ @@ -957,7 +978,7 @@ def delete_expired_events(self, before: "datetime", app_name: "str | None" = Non try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -974,7 +995,7 @@ def delete_idle_sessions(self, updated_before: "datetime", app_name: "str | None try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -991,7 +1012,7 @@ def delete_idle_user_states(self, updated_before: "datetime", app_name: "str | N try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -1003,7 +1024,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (app_name,)) + cur.execute(as_query(sql), (app_name,)) row = cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -1015,7 +1036,7 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (app_name, user_id)) + cur.execute(as_query(sql), (app_name, user_id)) row = cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -1029,7 +1050,7 @@ def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (app_name, Jsonb(state))) + cur.execute(as_query(sql), (app_name, Jsonb(state))) conn.commit() def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: @@ -1040,7 +1061,7 @@ def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]" """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (app_name, user_id, Jsonb(state))) + cur.execute(as_query(sql), (app_name, user_id, Jsonb(state))) conn.commit() def get_metadata(self, key: str) -> "str | None": @@ -1049,7 +1070,7 @@ def get_metadata(self, key: str) -> "str | None": try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (key,)) + cur.execute(as_query(sql), (key,)) row = cur.fetchone() return row["value"] if row is not None else None except errors.UndefinedTable: @@ -1063,7 +1084,7 @@ def set_metadata(self, key: str, value: str) -> None: """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (key, value)) + cur.execute(as_query(sql), (key, value)) conn.commit() def _sessions_table_ddl(self) -> str: @@ -1156,7 +1177,7 @@ def _insert_event(self, event_record: StoredEvent) -> None: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute( - sql.encode(), + as_query(sql), ( event_record["id"], event_record["app_name"], @@ -1275,7 +1296,7 @@ async def search_entries( try: async with self._config.provide_connection() as conn, conn.cursor() as cur: - await cur.execute(sql.encode(), params) + await cur.execute(as_query(sql), params) rows = await cur.fetchall() columns = [col[0] for col in cur.description or []] except errors.UndefinedTable: @@ -1295,7 +1316,7 @@ async def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = %s" async with self._config.provide_connection() as conn, conn.cursor() as cur: - await cur.execute(sql.encode(), (session_id,)) + await cur.execute(as_query(sql), (session_id,)) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1317,7 +1338,7 @@ 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, conn.cursor() as cur: - await cur.execute(sql.encode(), tuple(params) if params else None) + await cur.execute(as_query(sql), tuple(params) if params else None) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1459,7 +1480,7 @@ def search_entries( try: with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(sql.encode(), params) + cur.execute(as_query(sql), params) rows = cur.fetchall() columns = [col[0] for col in cur.description or []] except errors.UndefinedTable: @@ -1480,7 +1501,7 @@ def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = %s" with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(sql.encode(), (session_id,)) + cur.execute(as_query(sql), (session_id,)) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1501,7 +1522,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, conn.cursor() as cur: - cur.execute(sql.encode(), tuple(params) if params else None) + cur.execute(as_query(sql), tuple(params) if params else None) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 diff --git a/sqlspec/adapters/cockroach_psycopg/config.py b/sqlspec/adapters/cockroach_psycopg/config.py index 58b9c1ac5..6e04ebcef 100644 --- a/sqlspec/adapters/cockroach_psycopg/config.py +++ b/sqlspec/adapters/cockroach_psycopg/config.py @@ -17,6 +17,7 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_crdb as psycopg_crdb from sqlspec.adapters.cockroach_psycopg.core import ( apply_driver_features, + build_connection_config, build_statement_config, validate_follower_read_staleness, ) @@ -37,10 +38,9 @@ ) from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints -from sqlspec.utils.config_tools import normalize_connection_config if TYPE_CHECKING: - from collections.abc import Awaitable, Callable, Mapping + from collections.abc import Awaitable, Callable from types import TracebackType from sqlspec.core import StatementConfig @@ -70,6 +70,11 @@ class CockroachPsycopgConnectionConfig(TypedDict): connect_timeout: NotRequired[int] options: NotRequired[str] application_name: NotRequired[str] + gateway_region: NotRequired[str] + default_transaction_use_follower_reads: NotRequired[bool] + results_buffer_size: NotRequired[int] + statement_timeout: NotRequired[int] + idle_in_transaction_session_timeout: NotRequired[int] sslmode: NotRequired[str] sslcert: NotRequired[str] sslkey: NotRequired[str] @@ -142,33 +147,6 @@ class CockroachPsycopgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["poll_queue"]] -def build_connection_config( - connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", -) -> dict[str, Any]: - """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility. - - Maps connection string aliases (dsn, url, connection_string) to conninfo, database aliases - (database, db) to dbname, and user aliases (username) to user, while discarding redundant keys - that libpq rejects. - """ - config = normalize_connection_config(connection_config) - conninfo = ( - config.pop("conninfo", None) - or config.pop("dsn", None) - or config.pop("url", None) - or config.pop("connection_string", None) - ) - if conninfo is not None: - config["conninfo"] = conninfo - dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) - if dbname is not None: - config["dbname"] = dbname - user = config.pop("user", None) or config.pop("username", None) - if user is not None: - config["user"] = user - return config - - class CockroachPsycopgSyncConnectionContext(SyncPoolConnectionContext): """Context manager for CockroachDB psycopg connections.""" @@ -277,7 +255,7 @@ def _create_pool(self) -> "ConnectionPool": "name": all_config.pop("name", None), "timeout": all_config.pop("timeout", 30.0), "max_waiting": all_config.pop("max_waiting", 0), - "max_lifetime": all_config.pop("max_lifetime", 3600.0), + "max_lifetime": all_config.pop("max_lifetime", 1800.0), "max_idle": all_config.pop("max_idle", 600.0), "reconnect_timeout": all_config.pop("reconnect_timeout", 300.0), "reconnect_failed": all_config.pop("reconnect_failed", None), @@ -307,7 +285,6 @@ def _configure_connection(self, conn: "CockroachSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting - # Call user-provided callback after internal setup if self._user_connection_hook is not None: self._user_connection_hook(conn) @@ -504,7 +481,7 @@ async def _create_pool(self) -> "AsyncConnectionPool": "name": all_config.pop("name", None), "timeout": all_config.pop("timeout", 30.0), "max_waiting": all_config.pop("max_waiting", 0), - "max_lifetime": all_config.pop("max_lifetime", 3600.0), + "max_lifetime": all_config.pop("max_lifetime", 1800.0), "max_idle": all_config.pop("max_idle", 600.0), "reconnect_timeout": all_config.pop("reconnect_timeout", 300.0), "reconnect_failed": all_config.pop("reconnect_failed", None), @@ -541,7 +518,6 @@ async def _configure_async_connection(self, conn: "CockroachAsyncConnection") -> if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) - # 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/cockroach_psycopg/core.py b/sqlspec/adapters/cockroach_psycopg/core.py index 19dbc8257..e873ac1f7 100644 --- a/sqlspec/adapters/cockroach_psycopg/core.py +++ b/sqlspec/adapters/cockroach_psycopg/core.py @@ -1,26 +1,31 @@ """CockroachDB psycopg adapter compiled helpers.""" -import random import re +import secrets from typing import TYPE_CHECKING, Any, Final, cast from mypy_extensions import mypyc_attr from sqlglot import tokenize from sqlglot.tokenizer_core import TokenType +from typing_extensions import LiteralString from sqlspec.adapters.psycopg.core import apply_driver_features, build_statement_config, driver_profile from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError +from sqlspec.utils.config_tools import normalize_connection_config from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate if TYPE_CHECKING: from collections.abc import Mapping + from sqlspec.adapters.cockroach_psycopg.config import CockroachPsycopgPoolConfig from sqlspec.storage import StorageTelemetry __all__ = ( "CockroachPsycopgRetryConfig", "apply_driver_features", + "as_query", + "build_connection_config", "build_native_export", "build_native_import", "build_statement_config", @@ -33,14 +38,14 @@ "validate_follower_read_staleness", ) -# Retry configuration defaults (module-level for mypyc compatibility) _DEFAULT_MAX_RETRIES: Final[int] = 10 _DEFAULT_BASE_DELAY_MS: Final[float] = 50.0 _DEFAULT_MAX_DELAY_MS: Final[float] = 5000.0 _DEFAULT_ENABLE_LOGGING: Final[bool] = True +_MAX_EXCEPTION_CHAIN_DEPTH: Final[int] = 16 +_RNG: Final = secrets.SystemRandom() -# Keep this in sync with cockroach_asyncpg.core.CockroachAsyncpgRetryConfig. @mypyc_attr(allow_interpreted_subclasses=False) class CockroachPsycopgRetryConfig: """CockroachDB psycopg transaction retry configuration.""" @@ -70,6 +75,49 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachPsycop ) +def build_connection_config( + connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", +) -> dict[str, Any]: + """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility.""" + config = normalize_connection_config(connection_config) + conninfo = ( + config.pop("conninfo", None) + or config.pop("dsn", None) + or config.pop("url", None) + or config.pop("connection_string", None) + ) + if conninfo is not None: + config["conninfo"] = conninfo + dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) + if dbname is not None: + config["dbname"] = dbname + user = config.pop("user", None) or config.pop("username", None) + if user is not None: + config["user"] = user + + session_options: list[str] = [] + if "gateway_region" in config: + session_options.append(f"-c gateway_region={config.pop('gateway_region')}") + if "default_transaction_use_follower_reads" in config: + val = config.pop("default_transaction_use_follower_reads") + val_str = "on" if val is True or str(val).lower() == "on" else "off" + session_options.append(f"-c default_transaction_use_follower_reads={val_str}") + if "results_buffer_size" in config: + session_options.append(f"-c results_buffer_size={config.pop('results_buffer_size')}") + if "statement_timeout" in config: + session_options.append(f"-c statement_timeout={config.pop('statement_timeout')}") + if "idle_in_transaction_session_timeout" in config: + session_options.append( + f"-c idle_in_transaction_session_timeout={config.pop('idle_in_transaction_session_timeout')}" + ) + if session_options: + existing = config.get("options") + opt_str = " ".join(session_options) + config["options"] = f"{existing} {opt_str}" if existing else opt_str + + return config + + def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. @@ -90,17 +138,17 @@ def is_retryable_error(error: BaseException) -> bool: Returns: True when the transaction should be retried. """ - seen: set[int] = set() + depth = 0 current: BaseException | None = error - while current is not None and id(current) not in seen: - seen.add(id(current)) + while current is not None and depth < _MAX_EXCEPTION_CHAIN_DEPTH: if isinstance(current, SerializationConflictError): return True if has_sqlstate(current) and str(current.sqlstate) == "40001": return True if not isinstance(current, SQLSpecError): return False - current = cast("BaseException | None", cast("Any", current).__cause__) + current = cast("BaseException | None", getattr(current, "__cause__", None)) + depth += 1 return False @@ -114,7 +162,19 @@ def calculate_backoff_seconds(attempt: int, config: "CockroachPsycopgRetryConfig capped_ms: float = min(config.base_delay_ms * (2**attempt), config.max_delay_ms) if capped_ms <= 0.0: return 0.0 - return random.uniform(capped_ms / 2.0, capped_ms) / 1000.0 # noqa: S311 + return _RNG.uniform(capped_ms / 2.0, capped_ms) / 1000.0 + + +def as_query(sql: object) -> LiteralString: + """Prepare a SQL string for psycopg query execution without byte encoding. + + Args: + sql: The raw SQL query string or object. + + Returns: + The SQL query string typed as a LiteralString for driver query dispatch. + """ + return cast("LiteralString", sql) _STALENESS_LITERAL: Final[re.Pattern[str]] = re.compile(r"'[^'\\;]+'") diff --git a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py index b9a8a78ad..74e0ba79f 100644 --- a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py +++ b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py @@ -5,6 +5,7 @@ from mypy_extensions import mypyc_attr +from sqlspec.adapters.psycopg.data_dictionary import PsycopgAsyncDataDictionary, PsycopgSyncDataDictionary from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -29,10 +30,9 @@ unsupported_system_metadata_capability, ) from sqlspec.data_dictionary.dialects.cockroachdb import resolve_cockroachdb_json_type -from sqlspec.driver import AsyncDataDictionaryBase, SyncDataDictionaryBase if TYPE_CHECKING: - from sqlspec.adapters.cockroach_psycopg.driver import CockroachPsycopgAsyncDriver, CockroachPsycopgSyncDriver + from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver from sqlspec.core import SQL __all__ = ("CockroachPsycopgAsyncDataDictionary", "CockroachPsycopgSyncDataDictionary") @@ -58,7 +58,7 @@ @mypyc_attr(allow_interpreted_subclasses=True, native_class=False) -class CockroachPsycopgSyncDataDictionary(SyncDataDictionaryBase): +class CockroachPsycopgSyncDataDictionary(PsycopgSyncDataDictionary): """CockroachDB sync data dictionary.""" dialect: ClassVar[str] = "cockroachdb" @@ -67,13 +67,13 @@ def __init__(self) -> None: super().__init__() def get_metadata_capabilities( - self, driver: "CockroachPsycopgSyncDriver", domains: Sequence[str] | None = None + self, driver: "PsycopgSyncDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) def get_system_metadata_capabilities( - self, driver: "CockroachPsycopgSyncDriver", domains: Sequence[str] | None = None + self, driver: "PsycopgSyncDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -81,7 +81,7 @@ def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) def _select_domain( - self, driver: "CockroachPsycopgSyncDriver", domain: str, query_name: str, **parameters: Any + self, driver: "PsycopgSyncDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -89,17 +89,15 @@ def _select_domain( rows = driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - def get_schemas(self, driver: "CockroachPsycopgSyncDriver") -> MetadataResult: + def get_schemas(self, driver: "PsycopgSyncDriver") -> MetadataResult: """Get schema metadata.""" return self._select_domain(driver, "schemas", "list") - def get_objects(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_objects(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) - def get_table_details( - self, driver: "CockroachPsycopgSyncDriver", table: str, schema: str | None = None - ) -> MetadataResult: + def get_table_details(self, driver: "PsycopgSyncDriver", table: str, schema: str | None = None) -> MetadataResult: """Get rich table metadata.""" return self._select_domain( driver, @@ -110,7 +108,7 @@ def get_table_details( ) def get_constraints( - self, driver: "CockroachPsycopgSyncDriver", table: str | None = None, schema: str | None = None + self, driver: "PsycopgSyncDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -118,16 +116,16 @@ def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - def get_views(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_views(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - def get_routines(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_routines(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) def get_privileges( - self, driver: "CockroachPsycopgSyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "PsycopgSyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -136,7 +134,7 @@ def get_privileges( ) def get_dependencies( - self, driver: "CockroachPsycopgSyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "PsycopgSyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -146,7 +144,7 @@ def get_dependencies( def get_ddl( self, - driver: "CockroachPsycopgSyncDriver", + driver: "PsycopgSyncDriver", object_name: str, schema: str | None = None, *, @@ -172,7 +170,7 @@ def get_ddl( ) def get_system_metadata( - self, driver: "CockroachPsycopgSyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "PsycopgSyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -180,41 +178,41 @@ def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - def get_version(self, driver: "CockroachPsycopgSyncDriver") -> "VersionInfo | None": + def get_version(self, driver: "PsycopgSyncDriver") -> "VersionInfo | None": """Get CockroachDB version information.""" - driver_id = id(driver) - if driver_id in self._version_fetch_attempted: - return self._version_cache.get(driver_id) + cache_key = id(driver.connection) + if cache_key in self._version_fetch_attempted: + return self._version_cache.get(cache_key) version_value = driver.select_value_or_none(self.get_query("version", "current")) if not version_value: self._log_version_unavailable(type(self).dialect, "missing") - self.cache_version(driver_id, None) + self.cache_version(cache_key, None) return None version_info = self.parse_version_with_pattern(self.get_dialect_config().version_pattern, str(version_value)) if version_info is None: self._log_version_unavailable(type(self).dialect, "parse_failed") - self.cache_version(driver_id, None) + self.cache_version(cache_key, None) return None self._log_version_detected(type(self).dialect, version_info) - self.cache_version(driver_id, version_info) + self.cache_version(cache_key, version_info) return version_info - def get_feature_flag(self, driver: "CockroachPsycopgSyncDriver", feature: str) -> bool: + def get_feature_flag(self, driver: "PsycopgSyncDriver", feature: str) -> bool: """Check if CockroachDB supports a specific feature.""" version_info = self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - def get_optimal_type(self, driver: "CockroachPsycopgSyncDriver", type_category: str) -> str: + def get_optimal_type(self, driver: "PsycopgSyncDriver", type_category: str) -> str: """Get optimal CockroachDB type for a category.""" config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(self.get_version(driver)) return config.get_optimal_type(type_category) - def get_tables(self, driver: "CockroachPsycopgSyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": + def get_tables(self, driver: "PsycopgSyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": """Get tables sorted by dependency order.""" schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -226,7 +224,7 @@ def get_tables(self, driver: "CockroachPsycopgSyncDriver", schema: "str | None" ) def get_columns( - self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": """Get column information for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -250,7 +248,7 @@ def get_columns( ) def get_indexes( - self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": """Get index metadata for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -273,7 +271,7 @@ def get_indexes( ) def get_foreign_keys( - self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": """Get foreign key metadata.""" schema_name = self.resolve_schema(schema) @@ -294,7 +292,7 @@ def get_foreign_keys( @mypyc_attr(allow_interpreted_subclasses=True, native_class=False) -class CockroachPsycopgAsyncDataDictionary(AsyncDataDictionaryBase): +class CockroachPsycopgAsyncDataDictionary(PsycopgAsyncDataDictionary): """CockroachDB async data dictionary.""" dialect: ClassVar[str] = "cockroachdb" @@ -303,13 +301,13 @@ def __init__(self) -> None: super().__init__() async def get_metadata_capabilities( - self, driver: "CockroachPsycopgAsyncDriver", domains: Sequence[str] | None = None + self, driver: "PsycopgAsyncDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) async def get_system_metadata_capabilities( - self, driver: "CockroachPsycopgAsyncDriver", domains: Sequence[str] | None = None + self, driver: "PsycopgAsyncDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -317,7 +315,7 @@ async def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) async def _select_domain( - self, driver: "CockroachPsycopgAsyncDriver", domain: str, query_name: str, **parameters: Any + self, driver: "PsycopgAsyncDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -325,16 +323,16 @@ async def _select_domain( rows = await driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - async def get_schemas(self, driver: "CockroachPsycopgAsyncDriver") -> MetadataResult: + async def get_schemas(self, driver: "PsycopgAsyncDriver") -> MetadataResult: """Get schema metadata.""" return await self._select_domain(driver, "schemas", "list") - async def get_objects(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_objects(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return await self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) async def get_table_details( - self, driver: "CockroachPsycopgAsyncDriver", table: str, schema: str | None = None + self, driver: "PsycopgAsyncDriver", table: str, schema: str | None = None ) -> MetadataResult: """Get rich table metadata.""" return await self._select_domain( @@ -346,7 +344,7 @@ async def get_table_details( ) async def get_constraints( - self, driver: "CockroachPsycopgAsyncDriver", table: str | None = None, schema: str | None = None + self, driver: "PsycopgAsyncDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -354,16 +352,16 @@ async def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - async def get_views(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_views(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return await self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_routines(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_routines(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) async def get_privileges( - self, driver: "CockroachPsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "PsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -372,7 +370,7 @@ async def get_privileges( ) async def get_dependencies( - self, driver: "CockroachPsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "PsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -382,7 +380,7 @@ async def get_dependencies( async def get_ddl( self, - driver: "CockroachPsycopgAsyncDriver", + driver: "PsycopgAsyncDriver", object_name: str, schema: str | None = None, *, @@ -408,7 +406,7 @@ async def get_ddl( ) async def get_system_metadata( - self, driver: "CockroachPsycopgAsyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "PsycopgAsyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -416,43 +414,41 @@ async def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - async def get_version(self, driver: "CockroachPsycopgAsyncDriver") -> "VersionInfo | None": + async def get_version(self, driver: "PsycopgAsyncDriver") -> "VersionInfo | None": """Get CockroachDB version information.""" - driver_id = id(driver) - if driver_id in self._version_fetch_attempted: - return self._version_cache.get(driver_id) + cache_key = id(driver.connection) + if cache_key in self._version_fetch_attempted: + return self._version_cache.get(cache_key) version_value = await driver.select_value_or_none(self.get_query("version", "current")) if not version_value: self._log_version_unavailable(type(self).dialect, "missing") - self.cache_version(driver_id, None) + self.cache_version(cache_key, None) return None version_info = self.parse_version_with_pattern(self.get_dialect_config().version_pattern, str(version_value)) if version_info is None: self._log_version_unavailable(type(self).dialect, "parse_failed") - self.cache_version(driver_id, None) + self.cache_version(cache_key, None) return None self._log_version_detected(type(self).dialect, version_info) - self.cache_version(driver_id, version_info) + self.cache_version(cache_key, version_info) return version_info - async def get_feature_flag(self, driver: "CockroachPsycopgAsyncDriver", feature: str) -> bool: + async def get_feature_flag(self, driver: "PsycopgAsyncDriver", feature: str) -> bool: """Check if CockroachDB supports a specific feature.""" version_info = await self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - async def get_optimal_type(self, driver: "CockroachPsycopgAsyncDriver", type_category: str) -> str: + async def get_optimal_type(self, driver: "PsycopgAsyncDriver", type_category: str) -> str: """Get optimal CockroachDB type for a category.""" config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(await self.get_version(driver)) return config.get_optimal_type(type_category) - async def get_tables( - self, driver: "CockroachPsycopgAsyncDriver", schema: "str | None" = None - ) -> "list[TableMetadata]": + async def get_tables(self, driver: "PsycopgAsyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": """Get tables sorted by dependency order.""" schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -464,7 +460,7 @@ async def get_tables( ) async def get_columns( - self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": """Get column information for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -488,7 +484,7 @@ async def get_columns( ) async def get_indexes( - self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": """Get index metadata for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -511,7 +507,7 @@ async def get_indexes( ) async def get_foreign_keys( - self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": """Get foreign key metadata.""" schema_name = self.resolve_schema(schema) diff --git a/sqlspec/adapters/cockroach_psycopg/driver.py b/sqlspec/adapters/cockroach_psycopg/driver.py index 7871e1527..cd4be1209 100644 --- a/sqlspec/adapters/cockroach_psycopg/driver.py +++ b/sqlspec/adapters/cockroach_psycopg/driver.py @@ -1,6 +1,7 @@ """CockroachDB psycopg driver implementation.""" import asyncio +import contextlib import time from typing import TYPE_CHECKING, Any, TypeVar, cast @@ -14,6 +15,7 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_module as psycopg from sqlspec.adapters.cockroach_psycopg.core import ( CockroachPsycopgRetryConfig, + as_query, build_native_export, build_native_import, build_statement_config, @@ -30,9 +32,13 @@ CockroachPsycopgSyncDataDictionary, ) from sqlspec.adapters.psycopg.core import create_mapped_exception -from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver +from sqlspec.adapters.psycopg.driver import ( + PsycopgAsyncDriver, + PsycopgAsyncExceptionHandler, + PsycopgSyncDriver, + PsycopgSyncExceptionHandler, +) from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile -from sqlspec.driver import BaseAsyncExceptionHandler, BaseSyncExceptionHandler from sqlspec.utils.logging import get_logger if TYPE_CHECKING: @@ -55,7 +61,7 @@ _T = TypeVar("_T") -class CockroachPsycopgSyncExceptionHandler(BaseSyncExceptionHandler): +class CockroachPsycopgSyncExceptionHandler(PsycopgSyncExceptionHandler): """Context manager for handling CockroachDB psycopg exceptions.""" __slots__ = () @@ -69,7 +75,7 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False -class CockroachPsycopgAsyncExceptionHandler(BaseAsyncExceptionHandler): +class CockroachPsycopgAsyncExceptionHandler(PsycopgAsyncExceptionHandler): """Async context manager for handling CockroachDB psycopg exceptions.""" __slots__ = () @@ -105,7 +111,6 @@ 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 select_to_storage( @@ -208,7 +213,7 @@ def _execute_native_storage(self, command: str, parameters: "list[Any]") -> "lis rows = [] with self.with_cursor(self.connection) as cursor, handler: cursor.row_factory = dict_row - cursor.execute(command.encode("utf-8"), parameters) + cursor.execute(as_query(command), parameters) rows = cursor.fetchall() if handler.pending_exception is not None: raise handler.pending_exception @@ -261,58 +266,35 @@ def run_transaction_with_retry(self, operation: "Callable[[], _T]") -> _T: attempt += 1 def dispatch_execute(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - return self._dispatch_execute_impl(cursor, statement) - - def dispatch_execute_many(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - return self._dispatch_execute_many_impl(cursor, statement) - - def dispatch_execute_script(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - return self._dispatch_execute_script_impl(cursor, statement) - - def handle_database_exceptions(self) -> "CockroachPsycopgSyncExceptionHandler": # type: ignore[override] + opened_txn = False + if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): + self.begin() + opened_txn = True + try: + return super().dispatch_execute(cursor, statement) + finally: + if opened_txn: + with contextlib.suppress(Exception): + self.commit() + + def handle_database_exceptions(self) -> "CockroachPsycopgSyncExceptionHandler": return CockroachPsycopgSyncExceptionHandler() @property - def data_dictionary(self) -> "CockroachPsycopgSyncDataDictionary": # type: ignore[override] + def data_dictionary(self) -> "CockroachPsycopgSyncDataDictionary": if self._data_dictionary is None: - # Intentionally assign CockroachDB-specific data dictionary to parent slot - self._data_dictionary = CockroachPsycopgSyncDataDictionary() # type: ignore[assignment] + self._data_dictionary = CockroachPsycopgSyncDataDictionary() return cast("CockroachPsycopgSyncDataDictionary", self._data_dictionary) + def _follower_reads_enabled(self) -> bool: + return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) + def _apply_follower_reads(self) -> None: - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: + if not self._follower_reads_enabled() or not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) self.connection.execute(cast("Any", f"SET TRANSACTION AS OF SYSTEM TIME {staleness}")).close() - def _begin_follower_read_transaction(self) -> None: - """Open the transaction a follower read needs so the staleness clause can lead it. - - psycopg opens a transaction on the first statement, which would leave the - clause with nowhere to go, so a read opens one here when the caller has - not already done so. - """ - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: - return - if self._connection_in_transaction(): - return - self.begin() - - def _dispatch_execute_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - if statement.returns_rows(): - self._begin_follower_read_transaction() - return super().dispatch_execute(cursor, statement) - - def _dispatch_execute_many_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - return PsycopgSyncDriver.dispatch_execute_many(self, cursor, statement) - - def _dispatch_execute_script_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - return PsycopgSyncDriver.dispatch_execute_script(self, cursor, statement) - class CockroachPsycopgAsyncDriver(PsycopgAsyncDriver): """CockroachDB async driver using psycopg.crdb.""" @@ -336,7 +318,6 @@ 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 select_to_storage( @@ -439,7 +420,7 @@ async def _execute_native_storage(self, command: str, parameters: "list[Any]") - rows = [] async with self.with_cursor(self.connection) as cursor, handler: cursor.row_factory = dict_row - await cursor.execute(command.encode("utf-8"), parameters) + await cursor.execute(as_query(command), parameters) rows = await cursor.fetchall() if handler.pending_exception is not None: raise handler.pending_exception @@ -492,58 +473,35 @@ async def run_transaction_with_retry(self, operation: "Callable[[], Awaitable[_T attempt += 1 async def dispatch_execute(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_impl(cursor, statement) - - async def dispatch_execute_many(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_many_impl(cursor, statement) - - async def dispatch_execute_script(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - return await self._dispatch_execute_script_impl(cursor, statement) - - def handle_database_exceptions(self) -> "CockroachPsycopgAsyncExceptionHandler": # type: ignore[override] + opened_txn = False + if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): + await self.begin() + opened_txn = True + try: + return await super().dispatch_execute(cursor, statement) + finally: + if opened_txn: + with contextlib.suppress(Exception): + await self.commit() + + def handle_database_exceptions(self) -> "CockroachPsycopgAsyncExceptionHandler": return CockroachPsycopgAsyncExceptionHandler() @property - def data_dictionary(self) -> "CockroachPsycopgAsyncDataDictionary": # type: ignore[override] + def data_dictionary(self) -> "CockroachPsycopgAsyncDataDictionary": if self._data_dictionary is None: - # Intentionally assign CockroachDB-specific data dictionary to parent slot - self._data_dictionary = CockroachPsycopgAsyncDataDictionary() # type: ignore[assignment] + self._data_dictionary = CockroachPsycopgAsyncDataDictionary() return cast("CockroachPsycopgAsyncDataDictionary", self._data_dictionary) + def _follower_reads_enabled(self) -> bool: + return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) + async def _apply_follower_reads(self) -> None: - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: + if not self._follower_reads_enabled() or not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) cursor = await self.connection.execute(cast("Any", f"SET TRANSACTION AS OF SYSTEM TIME {staleness}")) await cursor.close() - async def _begin_follower_read_transaction(self) -> None: - """Open the transaction a follower read needs so the staleness clause can lead it. - - psycopg opens a transaction on the first statement, which would leave the - clause with nowhere to go, so a read opens one here when the caller has - not already done so. - """ - if not self.driver_features.get("enable_follower_reads", False): - return - if not self._follower_staleness: - return - if self._connection_in_transaction(): - return - await self.begin() - - async def _dispatch_execute_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - if statement.returns_rows(): - await self._begin_follower_read_transaction() - return await super().dispatch_execute(cursor, statement) - - async def _dispatch_execute_many_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - return await PsycopgAsyncDriver.dispatch_execute_many(self, cursor, statement) - - async def _dispatch_execute_script_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - return await PsycopgAsyncDriver.dispatch_execute_script(self, cursor, statement) - register_driver_profile("cockroach_psycopg", driver_profile) diff --git a/sqlspec/adapters/cockroach_psycopg/litestar/store.py b/sqlspec/adapters/cockroach_psycopg/litestar/store.py index 8b915ca1e..c7bbad239 100644 --- a/sqlspec/adapters/cockroach_psycopg/litestar/store.py +++ b/sqlspec/adapters/cockroach_psycopg/litestar/store.py @@ -6,6 +6,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_dict_row as dict_row +from sqlspec.adapters.cockroach_psycopg.core import as_query from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -67,7 +68,7 @@ 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: async with conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), (key,)) + await cur.execute(as_query(sql), (key,)) row = await cur.fetchone() if row is None: @@ -81,7 +82,7 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by SET expires_at = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ - await conn.execute(update_sql.encode(), (new_expires_at, key)) + await conn.execute(as_query(update_sql), (new_expires_at, key)) await conn.commit() return bytes(row["data"]) @@ -102,7 +103,7 @@ 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: - await conn.execute(sql.encode(), (key, data, expires_at)) + await conn.execute(as_query(sql), (key, data, expires_at)) await conn.commit() async def delete(self, key: str) -> None: @@ -110,7 +111,7 @@ async def delete(self, key: str) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - await conn.execute(sql.encode(), (key,)) + await conn.execute(as_query(sql), (key,)) await conn.commit() async def delete_all(self) -> None: @@ -118,7 +119,7 @@ async def delete_all(self) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - await conn.execute(sql.encode()) + await conn.execute(as_query(sql)) await conn.commit() self._log_delete_all() @@ -131,7 +132,7 @@ async def exists(self, key: str) -> bool: conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: - await cur.execute(sql.encode(), (key,)) + await cur.execute(as_query(sql), (key,)) row = await cur.fetchone() return row is not None @@ -144,7 +145,7 @@ async def expires_in(self, key: str) -> "int | None": conn_context = self._config.provide_connection() async with conn_context as conn: async with conn.cursor(row_factory=dict_row) as cur: - await cur.execute(sql.encode(), (key,)) + await cur.execute(as_query(sql), (key,)) row = await cur.fetchone() if row is None or row["expires_at"] is None: @@ -164,7 +165,7 @@ async def delete_expired(self) -> int: conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: - await cur.execute(sql.encode()) + await cur.execute(as_query(sql)) await conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 if count > 0: @@ -268,7 +269,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (key,)) + cur.execute(as_query(sql), (key,)) row = cur.fetchone() if row is None: @@ -282,7 +283,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | SET expires_at = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ - conn.execute(update_sql.encode(), (new_expires_at, key)) + conn.execute(as_query(update_sql), (new_expires_at, key)) conn.commit() return bytes(row["data"]) @@ -302,21 +303,21 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No """ with self._config.provide_connection() as conn: - conn.execute(sql.encode(), (key, data, expires_at)) + conn.execute(as_query(sql), (key, data, expires_at)) conn.commit() def _delete(self, key: str) -> None: sql = f"DELETE FROM {self._table_name} WHERE session_id = %s" with self._config.provide_connection() as conn: - conn.execute(sql.encode(), (key,)) + conn.execute(as_query(sql), (key,)) conn.commit() def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name}" with self._config.provide_connection() as conn: - conn.execute(sql.encode()) + conn.execute(as_query(sql)) conn.commit() self._log_delete_all() @@ -328,7 +329,7 @@ def _exists(self, key: str) -> bool: """ with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(sql.encode(), (key,)) + cur.execute(as_query(sql), (key,)) row = cur.fetchone() return row is not None @@ -340,7 +341,7 @@ def _expires_in(self, key: str) -> "int | None": with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: - cur.execute(sql.encode(), (key,)) + cur.execute(as_query(sql), (key,)) row = cur.fetchone() if row is None or row["expires_at"] is None: @@ -359,7 +360,7 @@ def _delete_expired(self) -> int: sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP" with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(sql.encode()) + cur.execute(as_query(sql)) conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 if count > 0: diff --git a/sqlspec/adapters/psqlpy/adk/store.py b/sqlspec/adapters/psqlpy/adk/store.py index ecbe6df2f..41917eea9 100644 --- a/sqlspec/adapters/psqlpy/adk/store.py +++ b/sqlspec/adapters/psqlpy/adk/store.py @@ -73,6 +73,8 @@ class PsqlpyADKStore(BaseAsyncADKStore["PsqlpyConfig"]): __slots__ = () + _config: "PsqlpyConfig" + def __init__(self, config: "PsqlpyConfig") -> None: super().__init__(config) @@ -91,26 +93,38 @@ async def create_tables(self) -> None: async def create_session( self, session_id: str, app_name: str, user_id: str, state: "dict[str, Any]", owner_id: "Any | None" = None ) -> StoredSession: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: if self._owner_id_column_name: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, {self._owner_id_column_name}, state, create_time, update_time) VALUES ($1, $2, $3, $4, $5, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ - await conn.execute(sql, [session_id, app_name, user_id, owner_id, state]) + single_result = await conn.fetch_row(sql, [session_id, app_name, user_id, owner_id, state]) else: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time) VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) + RETURNING id, app_name, user_id, state, create_time, update_time """ - await conn.execute(sql, [session_id, app_name, user_id, state]) + single_result = await conn.fetch_row(sql, [session_id, app_name, user_id, state]) - res = await self.get_session(app_name, user_id, session_id) - if res is None: - msg = "Failed to retrieve created session." + if not single_result: + msg = "Failed to fetch created session" + raise RuntimeError(msg) + row = single_result.result() + if not row: + msg = "Failed to fetch created session" raise RuntimeError(msg) - return res + return StoredSession( + id=row["id"], + app_name=row["app_name"], + user_id=row["user_id"], + state=row["state"], + create_time=row["create_time"], + update_time=row["update_time"], + ) async def get_session( self, app_name: str, user_id: str, session_id: str, *, renew_for: "int | timedelta | None" = None @@ -130,14 +144,14 @@ async def get_session( """ try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] - result = await conn.fetch(sql, [app_name, user_id, session_id]) - rows: list[dict[str, Any]] = result.result() if result else [] - - if not rows: + async with self._config.provide_connection() as conn: + single_result = await conn.fetch_row(sql, [app_name, user_id, session_id]) + if not single_result: + return None + row = single_result.result() + if not row: return None - row = rows[0] return StoredSession( id=row["id"], app_name=row["app_name"], @@ -158,7 +172,7 @@ async def update_session_state(self, app_name: str, user_id: str, session_id: st WHERE app_name = $2 AND user_id = $3 AND id = $4 """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute(sql, [state, app_name, user_id, session_id]) async def list_sessions( @@ -196,7 +210,7 @@ async def list_sessions( """ try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, params) rows: list[dict[str, Any]] = result.result() if result else [] @@ -219,7 +233,7 @@ async def list_sessions( async def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: sql = f"DELETE FROM {self._session_table} WHERE app_name = $1 AND user_id = $2 AND id = $3" - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute(sql, [app_name, user_id, session_id]) async def append_event(self, event_record: StoredEvent) -> None: @@ -229,7 +243,7 @@ async def append_event(self, event_record: StoredEvent) -> None: ) VALUES ($1, $2, $3, $4, $5, $6, $7) """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute( sql, [ @@ -280,7 +294,7 @@ async def append_event_and_update_state( update_time = CURRENT_TIMESTAMP """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: try: await conn.execute("BEGIN") await conn.execute( @@ -350,7 +364,7 @@ async def get_events( """ try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, params) rows: list[dict[str, Any]] = result.result() if result else [] @@ -382,7 +396,7 @@ async def delete_expired_events(self, before: "datetime", app_name: "str | None" params = [before] try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: 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 @@ -404,7 +418,7 @@ async def delete_idle_sessions(self, updated_before: "datetime", app_name: "str params = [updated_before] try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: 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 @@ -428,7 +442,7 @@ async def delete_idle_user_states(self, updated_before: "datetime", app_name: "s params = [updated_before] try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: 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 @@ -443,7 +457,7 @@ async def get_app_state(self, app_name: str) -> "dict[str, Any] | None": sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = $1" try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, [app_name]) rows: list[dict[str, Any]] = result.result() if result else [] return rows[0]["state"] if rows else None @@ -456,7 +470,7 @@ async def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | sql = f"SELECT state FROM {self._user_state_table} WHERE app_name = $1 AND user_id = $2" try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, [app_name, user_id]) rows: list[dict[str, Any]] = result.result() if result else [] return rows[0]["state"] if rows else None @@ -474,7 +488,7 @@ async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None update_time = CURRENT_TIMESTAMP """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute(sql, [app_name, state]) async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: @@ -486,14 +500,14 @@ async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, update_time = CURRENT_TIMESTAMP """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute(sql, [app_name, user_id, state]) async def get_metadata(self, key: str) -> "str | None": sql = f"SELECT value FROM {self._metadata_table} WHERE key = $1" try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, [key]) rows: list[dict[str, Any]] = result.result() if result else [] return rows[0]["value"] if rows else None @@ -509,7 +523,7 @@ async def set_metadata(self, key: str, value: str) -> None: ON CONFLICT (key) DO UPDATE SET value = EXCLUDED.value """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: await conn.execute(sql, [key, value]) async def _sessions_table_ddl(self) -> str: @@ -620,6 +634,8 @@ class PsqlpyADKMemoryStore(BaseAsyncADKMemoryStore["PsqlpyConfig"]): __slots__ = () + _config: "PsqlpyConfig" + def __init__(self, config: "PsqlpyConfig") -> None: """Initialize Psqlpy memory store.""" super().__init__(config) @@ -668,7 +684,7 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " ON CONFLICT (event_id) DO NOTHING """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: for entry in entries: if self._owner_id_column_name: params = [ @@ -727,7 +743,7 @@ async def search_entries( if self._use_fts: try: return await self._search_entries_fts(query, app_name, user_id, effective_limit) - except Exception as exc: # pragma: no cover + 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) except Exception as e: @@ -743,7 +759,7 @@ async def delete_entries_by_session(self, session_id: str) -> int: delete_sql = f"DELETE FROM {self._memory_table} WHERE session_id = $1" try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: count_result = await conn.fetch(count_sql, [session_id]) count_rows: list[dict[str, Any]] = count_result.result() if count_result else [] count = int(count_rows[0]["count"]) if count_rows else 0 @@ -770,7 +786,7 @@ async def delete_entries_older_than( """ try: - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: count_result = await conn.fetch(count_sql, []) count_rows: list[dict[str, Any]] = count_result.result() if count_result else [] count = int(count_rows[0]["count"]) if count_rows else 0 @@ -849,7 +865,7 @@ async def _search_entries_fts( ORDER BY rank DESC, timestamp DESC LIMIT {p_lim} """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, [*scope_params, limit]) rows: list[dict[str, Any]] = result.result() if result else [] return _rows_to_records(rows) @@ -883,7 +899,7 @@ async def _search_entries_simple( ORDER BY timestamp DESC LIMIT {p_lim} """ - async with self._config.provide_connection() as conn: # pyright: ignore[reportAttributeAccessIssue] + async with self._config.provide_connection() as conn: result = await conn.fetch(sql, [*scope_params, limit]) rows: list[dict[str, Any]] = result.result() if result else [] return _rows_to_records(rows) diff --git a/sqlspec/adapters/psqlpy/config.py b/sqlspec/adapters/psqlpy/config.py index 5f8a7e4c3..b5b74db3d 100644 --- a/sqlspec/adapters/psqlpy/config.py +++ b/sqlspec/adapters/psqlpy/config.py @@ -241,7 +241,6 @@ def __init__( self._user_connection_hook: Callable[[PsqlpyConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) - self._initialized_connection_ids: set[int] = set() self._pgvector_available: bool | None = None self._paradedb_available: bool | None = None self._pg_textsearch_available: bool | None = None @@ -284,12 +283,30 @@ async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: ) self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") - conn_id = id(connection) - if conn_id in self._initialized_connection_ids: + if getattr(connection, "_sqlspec_initialized", False): return if self._user_connection_hook is not None: await self._user_connection_hook(connection) - self._initialized_connection_ids.add(conn_id) + setattr(connection, "_sqlspec_initialized", True) + + def get_pool_status(self) -> "dict[str, int] | None": + """Return connection pool status metrics if pool is active.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "status"): + status = pool.status() + return { + "max_size": status.max_size, + "size": status.size, + "available": status.available, + "waiting": status.waiting, + } + return None + + def resize_pool(self, new_max_size: int) -> None: + """Dynamically resize the active connection pool.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "resize"): + pool.resize(new_max_size) async def _create_pool(self) -> "ConnectionPool": """Create the actual async connection pool.""" diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index f525a7c50..23762920c 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -77,6 +77,7 @@ "get_parameter_casts", "is_postgres_extension_active", "prepare_parameters_with_casts", + "records_to_arrow_table", "resolve_postgres_extension_state", "resolve_runtime_statement_config", "split_schema_and_table", @@ -91,6 +92,7 @@ "TIMESTAMP WITHOUT TIME ZONE", }) _UUID_CASTS: Final[frozenset[str]] = frozenset({"UUID"}) +_VECTOR_CASTS: Final[frozenset[str]] = frozenset({"VECTOR", "HALFVEC", "SPARSEVEC"}) _DECIMAL_NORMALIZER = build_nested_decimal_normalizer(mode="float") _JSONB_TYPE: type[Any] | None = None try: @@ -103,6 +105,34 @@ _DML_COUNT_CTE_ALIAS: Final = "_sqlspec_affected" _DML_COUNT_COLUMN: Final = "_sqlspec_rows_affected" _DML_COUNT_QUERY_CACHE_SIZE: Final = 1024 +_PSQLPY_ACCEPTED_POOL_KWARGS: Final[frozenset[str]] = frozenset({ + "dsn", + "username", + "password", + "host", + "hosts", + "port", + "ports", + "db_name", + "target_session_attrs", + "options", + "application_name", + "connect_timeout_sec", + "connect_timeout_nanosec", + "tcp_user_timeout_sec", + "tcp_user_timeout_nanosec", + "keepalives", + "keepalives_idle_sec", + "keepalives_idle_nanosec", + "keepalives_interval_sec", + "keepalives_interval_nanosec", + "keepalives_retries", + "load_balance_hosts", + "max_db_pool_size", + "conn_recycling_method", + "ssl_mode", + "ca_file", +}) logger = get_logger("sqlspec.adapters.psqlpy.core") _NUMERIC_COERCE_TYPES: "tuple[type[Any], ...]" = (float, decimal.Decimal, list, tuple, dict) @@ -149,7 +179,7 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str connection_config: Raw connection configuration mapping. Returns: - Dictionary with connection parameters. + Dictionary with sanitized connection parameters accepted by psqlpy. """ config = {key: value for key, value in connection_config.items() if value is not None} dsn = ( @@ -171,7 +201,32 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str username = config.pop("username", None) or config.pop("user", None) if username is not None: config["username"] = username - return config + max_size = config.pop("max_size", None) or config.pop("max_db_pool_size", None) + if max_size is not None: + config["max_db_pool_size"] = max_size + timeout = ( + config.pop("connect_timeout_sec", None) or config.pop("connect_timeout", None) or config.pop("timeout", None) + ) + if timeout is not None: + config["connect_timeout_sec"] = int(timeout) + + valid_config: dict[str, Any] = {} + extra_params: dict[str, Any] = {} + for key, value in config.items(): + if key in _PSQLPY_ACCEPTED_POOL_KWARGS: + valid_config[key] = value + else: + extra_params[key] = value + + if extra_params and "dsn" in valid_config: + dsn_val = str(valid_config["dsn"]) + if "?" in dsn_val: + query_suffix = "&" + "&".join(f"{k}={v}" for k, v in extra_params.items()) + valid_config["dsn"] = dsn_val + query_suffix + elif dsn_val.startswith(("postgresql://", "postgres://")): + query_suffix = "?" + "&".join(f"{k}={v}" for k, v in extra_params.items()) + valid_config["dsn"] = dsn_val + query_suffix + return valid_config def apply_driver_features( @@ -197,11 +252,12 @@ def apply_driver_features( return statement_config, features -def collect_rows(query_result: Any | None) -> "tuple[list[dict[str, Any]], list[str]]": +def collect_rows(query_result: Any | None, as_records: bool = True) -> "tuple[list[Any], list[str]]": """Collect psqlpy rows and column names. Args: query_result: Result returned from cursor.fetch(). + as_records: Whether to return Record objects if available. Returns: Tuple of (rows, column_names). @@ -209,12 +265,56 @@ def collect_rows(query_result: Any | None) -> "tuple[list[dict[str, Any]], list[ if not query_result: return [], [] + if as_records and hasattr(query_result, "records"): + records = cast("list[Any]", query_result.records()) + if not records: + return [], [] + first = records[0] + column_names = list(first.keys()) if hasattr(first, "keys") else [] + return records, column_names + dict_rows = cast("list[dict[str, Any]]", query_result if isinstance(query_result, list) else query_result.result()) if not dict_rows: return [], [] return dict_rows, list(dict_rows[0]) +def records_to_arrow_table(records: list[Any], columns: list[str], schema: Any = None) -> Any: + """Construct a pyarrow Table from records and column names using columnar arrays. + + Args: + records: List of records or row dictionaries. + columns: Column names corresponding to the records. + schema: Optional pyarrow schema. + + Returns: + A pyarrow Table. + """ + import pyarrow as pa + + if not records: + if schema is not None: + return pa.Table.from_batches([], schema=schema) + return pa.Table.from_arrays([pa.array([]) for _ in columns], names=columns) + + first = records[0] + is_dict = isinstance(first, dict) + if schema is not None: + arrays = [ + pa.array( + [r.get(columns[col_idx]) if is_dict else r[col_idx] for r in records], type=schema.field(col_idx).type + ) + for col_idx in range(len(columns)) + ] + return pa.Table.from_arrays(arrays, schema=schema) + + arrays = [ + pa.array([r.get(columns[col_idx]) if is_dict else r[col_idx] for r in records]) + for col_idx in range(len(columns)) + ] + return pa.Table.from_arrays(arrays, names=columns) + + class PsqlpyStreamSource: """Compiled async chunk source streaming dict rows from a psqlpy server-side cursor. @@ -224,6 +324,13 @@ class PsqlpyStreamSource: left untouched. """ + _chunk_size: int + _cursor: Any + _driver: Any + _parameters: Any + _sql: str + _transaction: Any + __slots__ = ("_chunk_size", "_cursor", "_driver", "_parameters", "_sql", "_transaction") def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: @@ -231,8 +338,8 @@ def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> N self._sql = sql self._parameters = parameters self._chunk_size = chunk_size - self._cursor: Any = None - self._transaction: Any = None + self._cursor = None + self._transaction = None async def start(self) -> None: handler = self._driver.handle_database_exceptions() @@ -258,12 +365,14 @@ async def _start(self) -> None: await transaction.rollback() raise - async def fetch_chunk(self) -> "list[dict[str, Any]]": + async def fetch_chunk(self) -> "list[Any]": handler = self._driver.handle_database_exceptions() query_result = await self._driver._run_with_exception_handler(handler, self._cursor.fetchmany, self._chunk_size) self._driver._check_pending_exception(handler) if query_result is None: return [] + if hasattr(query_result, "records"): + return cast("list[Any]", query_result.records()) return cast("list[dict[str, Any]]", query_result.result()) async def close(self, error: bool = False) -> None: @@ -286,6 +395,7 @@ async def close(self, error: bool = False) -> None: def coerce_numeric_for_write(value: Any) -> Any: + """Coerce numerical values to Decimal for precise Postgres numeric writes.""" if isinstance(value, float): return decimal.Decimal(str(value)) if isinstance(value, decimal.Decimal): @@ -598,6 +708,10 @@ def _coerce_parameter_for_cast(value: Any, cast_type: str, serializer: "Callable return _coerce_uuid_parameter(value) if upper_cast in _TIMESTAMP_CASTS: return _coerce_timestamp_parameter(value) + if upper_cast in _VECTOR_CASTS: + from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector + + return coerce_pgvector(value) return value diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index 69b9f88c0..850372ffe 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -4,8 +4,12 @@ and transaction management. """ +import contextlib +from time import perf_counter from typing import TYPE_CHECKING, Any, cast +from mypy_extensions import mypyc_attr + from sqlspec.adapters.psqlpy._typing import PsqlpyCursor, PsqlpyDatabaseError, PsqlpyError, PsqlpySessionContext from sqlspec.adapters.psqlpy.core import ( _DML_COUNT_COLUMN, @@ -22,12 +26,23 @@ format_table_identifier, get_parameter_casts, prepare_parameters_with_casts, + records_to_arrow_table, split_schema_and_table, ) from sqlspec.adapters.psqlpy.data_dictionary import PsqlpyDataDictionary -from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.core import ( + SQL, + StackResult, + StatementConfig, + create_arrow_result, + get_cache_config, + register_driver_profile, +) +from sqlspec.core.stack import StatementStack from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler +from sqlspec.driver._common import validate_savepoint_name from sqlspec.exceptions import SQLSpecError +from sqlspec.utils.schema import to_value_type from sqlspec.utils.text import normalize_identifier, quote_identifier if TYPE_CHECKING: @@ -64,6 +79,7 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsqlpyDriver(AsyncDriverAdapterBase): """PostgreSQL driver implementation using psqlpy. @@ -71,7 +87,11 @@ class PsqlpyDriver(AsyncDriverAdapterBase): and transaction management. """ - __slots__ = ("_data_dictionary", "_transaction_active") + _data_dictionary: PsqlpyDataDictionary | None + _transaction_active: bool + _json_columns_cache: dict[tuple[str | None, str], set[str]] + + __slots__ = ("_data_dictionary", "_json_columns_cache", "_transaction_active") dialect = "postgres" def __init__( @@ -86,8 +106,9 @@ def __init__( ) super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) - self._data_dictionary: PsqlpyDataDictionary | None = None + self._data_dictionary = None self._transaction_active = False + self._json_columns_cache = {} async def dispatch_execute(self, cursor: "PsqlpyConnection", statement: SQL) -> "ExecutionResult": """Execute single SQL statement. @@ -116,6 +137,18 @@ async def dispatch_execute(self, cursor: "PsqlpyConnection", statement: SQL) -> ) if statement.operation_type in {"INSERT", "UPDATE", "DELETE"}: + if "returning" in sql.lower(): + query_result = await cursor.fetch(sql, params) + dict_rows, column_names = collect_rows(query_result) + rows_affected = len(dict_rows) + return self.create_execution_result( + cursor, + selected_data=dict_rows, + column_names=column_names, + data_row_count=rows_affected, + rowcount_override=rows_affected, + is_select_result=statement.returns_rows(), + ) count_sql = _dml_count_query(sql) if count_sql is not None: count_result = await cursor.fetch(count_sql, params) @@ -158,7 +191,7 @@ async def dispatch_execute_many(self, cursor: "PsqlpyConnection", statement: SQL return self.create_execution_result(cursor, rowcount_override=rows_affected, is_many_result=True) async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: SQL) -> "ExecutionResult": - """Execute SQL script with statement splitting. + """Execute SQL script with statement splitting or batch execution. Args: cursor: Psqlpy connection object @@ -170,6 +203,18 @@ async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: S sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) prepared_parameters = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) statement_config = statement.statement_config + + if not prepared_parameters and hasattr(cursor, "execute_batch"): + statements = self.split_script_statements(sql, statement_config, strip_trailing_semicolon=True) + exc_handler = self.handle_database_exceptions() + async with exc_handler: + await cursor.execute_batch(sql) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + return self.create_execution_result( + cursor, statement_count=len(statements), successful_statements=len(statements), is_script_result=True + ) + statements = self.split_script_statements(sql, statement_config, strip_trailing_semicolon=True) successful_count = 0 @@ -216,6 +261,24 @@ async def rollback(self) -> None: finally: self._transaction_active = False + async def savepoint(self, name: str) -> None: + """Create a savepoint within the current transaction.""" + validate_savepoint_name(name) + quoted_name = quote_identifier(name) + await self.connection.execute(f"SAVEPOINT {quoted_name}") + + async def release_savepoint(self, name: str) -> None: + """Release a savepoint within the current transaction.""" + validate_savepoint_name(name) + quoted_name = quote_identifier(name) + await self.connection.execute(f"RELEASE SAVEPOINT {quoted_name}") + + async def rollback_savepoint(self, name: str) -> None: + """Rollback to a savepoint within the current transaction.""" + validate_savepoint_name(name) + quoted_name = quote_identifier(name) + await self.connection.execute(f"ROLLBACK TO SAVEPOINT {quoted_name}") + async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" normalized_schema = normalize_identifier(schema, "postgres") @@ -247,6 +310,10 @@ async def _resolve_json_columns(self, schema_name: "str | None", table_name: str Returns: Names of columns typed json or jsonb. """ + cache_key = (schema_name, table_name) + if cache_key in self._json_columns_cache: + return self._json_columns_cache[cache_key] + qualified = quote_identifier(table_name) if schema_name is not None: qualified = f"{quote_identifier(schema_name)}.{qualified}" @@ -261,7 +328,9 @@ async def _resolve_json_columns(self, schema_name: "str | None", table_name: str [qualified], ) data, _ = collect_rows(rows) - return {str(row["column_name"]) for row in data} + result = {str(row["column_name"]) for row in data} + self._json_columns_cache[cache_key] = result + return result async def has_schema(self, schema: str) -> bool: """Return whether a PostgreSQL schema exists.""" @@ -299,6 +368,173 @@ def handle_database_exceptions(self) -> "PsqlpyExceptionHandler": """ return PsqlpyExceptionHandler() + async def execute_stack( + self, stack: "StatementStack", *, continue_on_error: bool = False + ) -> "tuple[StackResult, ...]": + """Execute a StatementStack using psqlpy pipelining when available.""" + if not isinstance(stack, StatementStack) or not stack or self.stack_native_disabled or continue_on_error: + return await super().execute_stack(stack, continue_on_error=continue_on_error) + + queries: list[tuple[str, list[Any] | None]] = [] + prepared_operations: list[tuple[Any, Any]] = [] + + for operation in stack.operations: + kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else {} + config = kwargs.pop("statement_config", None) or self.statement_config + sql_statement = self.prepare_statement( + operation.statement, operation.arguments, statement_config=config, kwargs=kwargs + ) + if sql_statement.is_script or sql_statement.is_many: + return await super().execute_stack(stack, continue_on_error=continue_on_error) + sql, params = self._compiled_sql(sql_statement, config) + p_list = list(params) if isinstance(params, (list, tuple)) else None + queries.append((sql, p_list)) + prepared_operations.append((operation, sql_statement)) + + transaction = self.connection.transaction() + needs_commit = False + if not self._connection_in_transaction(): + await transaction.begin() + needs_commit = True + + results: list[StackResult] = [] + try: + query_results = await transaction.pipeline(queries) + if needs_commit: + await transaction.commit() + for (_op, stmt), q_res in zip(prepared_operations, query_results, strict=False): + rows, column_names = collect_rows(q_res) + exec_result = self.create_execution_result( + self.connection, + selected_data=rows, + column_names=column_names, + data_row_count=len(rows), + is_select_result=stmt.returns_rows(), + ) + sql_result = self.build_statement_result(stmt, exec_result) + results.append(StackResult(result=sql_result)) + except Exception as exc: + if needs_commit: + with contextlib.suppress(Exception): + await transaction.rollback() + msg = f"Pipelined stack execution failed: {exc}" + raise SQLSpecError(msg) from exc + + return tuple(results) + + async def select_to_arrow( + self, + statement: Any, + /, + *parameters: Any, + statement_config: "StatementConfig | None" = None, + return_format: str = "table", + native_only: bool = False, + batch_size: int | None = None, + arrow_schema: Any = None, + **kwargs: Any, + ) -> "ArrowResult": + """Execute a query and return results formatted as Apache Arrow.""" + import pyarrow as pa + + config = statement_config or self.statement_config + sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) + sql, prepared_parameters = self._compiled_sql(sql_statement, config) + params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] + + start_time = perf_counter() + query_result: Any = None + exc_handler = self.handle_database_exceptions() + async with exc_handler, self.with_cursor(self.connection) as cursor: + query_result = await cursor.fetch(sql, params) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + execution_time = perf_counter() - start_time + + records = query_result.records() if hasattr(query_result, "records") else query_result.result() + columns = list(records[0].keys()) if records and hasattr(records[0], "keys") else [] + + table = records_to_arrow_table(records, columns, schema=arrow_schema) + + if return_format == "table": + data: Any = table + elif return_format == "batch": + batches = table.to_batches() + data = batches[0] if batches else pa.RecordBatch.from_arrays([], schema=table.schema) + elif return_format == "batches": + data = table.to_batches(max_chunksize=batch_size) if batch_size else table.to_batches() + elif return_format == "reader": + data = table.to_reader(max_chunksize=batch_size) + else: + data = table + + return create_arrow_result( + statement=sql_statement, + data=data, + rows_affected=len(records), + execution_time=execution_time, + metadata={"columns": columns}, + ) + + async def select_one_or_none( + self, + statement: Any, + /, + *parameters: Any, + schema_type: Any = None, + statement_config: "StatementConfig | None" = None, + **kwargs: Any, + ) -> Any: + """Execute a query returning at most one row using fetch_row fast-path.""" + config = statement_config or self.statement_config + sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) + sql, prepared_parameters = self._compiled_sql(sql_statement, config) + params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] + + single_result: Any = None + exc_handler = self.handle_database_exceptions() + async with exc_handler, self.with_cursor(self.connection) as cursor: + single_result = await cursor.fetch_row(sql, params) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + + if single_result is None: + return None + row_dict = single_result.result() if hasattr(single_result, "result") else dict(cast("Any", single_result)) + if not row_dict: + return None + if schema_type is not None: + return self.to_schema(row_dict, schema_type=schema_type) + return row_dict + + async def select_value( + self, + statement: Any, + /, + *parameters: Any, + value_type: Any = None, + statement_config: "StatementConfig | None" = None, + **kwargs: Any, + ) -> Any: + """Execute a query returning a scalar value using fetch_val fast-path.""" + config = statement_config or self.statement_config + sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) + sql, prepared_parameters = self._compiled_sql(sql_statement, config) + params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] + + val: Any = None + exc_handler = self.handle_database_exceptions() + async with exc_handler, self.with_cursor(self.connection) as cursor: + val = await cursor.fetch_val(sql, params) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + + if val is None: + return None + if value_type is not None: + return to_value_type(val, value_type) + return val + async def select_to_storage( self, statement: "SQL | str", @@ -322,6 +558,72 @@ async def select_to_storage( self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload, telemetry) + async def load_from_records( + self, + table: str, + records: "Sequence[Mapping[str, Any]] | Sequence[Sequence[Any]]", + *, + columns: "list[str] | None" = None, + overwrite: bool = False, + partitioner: "dict[str, object] | None" = None, + telemetry: "StorageTelemetry | None" = None, + ) -> "StorageBridgeJob": + """Load Python records into PostgreSQL via psqlpy binary COPY.""" + self._require_capability("arrow_import_enabled") + if overwrite: + qualified = format_table_identifier(table) + exc_handler = self.handle_database_exceptions() + async with exc_handler, self.with_cursor(self.connection) as cursor: + await cursor.execute(f"TRUNCATE TABLE {qualified}") + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + + if not records: + empty_payload: StorageTelemetry = {"destination": table, "rows_processed": 0, "bytes_processed": 0} + self._attach_partition_telemetry(empty_payload, partitioner) + return self._storage_job(empty_payload, telemetry) + + schema_name, table_name = split_schema_and_table(table) + first_record = records[0] + from collections.abc import Mapping as MappingABC + + if columns is None: + if isinstance(first_record, MappingABC): + resolved_columns = list(first_record.keys()) + else: + msg = "columns must be provided when records are sequences" + raise SQLSpecError(msg) + else: + resolved_columns = columns + + if isinstance(first_record, MappingABC): + row_tuples = [ + tuple(r.get(col) for col in resolved_columns) for r in cast("Sequence[Mapping[str, Any]]", records) + ] + else: + row_tuples = [tuple(r) for r in cast("Sequence[Sequence[Any]]", records)] + + json_columns = await self._resolve_json_columns(schema_name, table_name) + coerced_records = coerce_json_columns(row_tuples, resolved_columns, json_columns) + + copy_kwargs: dict[str, Any] = {"columns": resolved_columns} + if schema_name: + copy_kwargs["schema_name"] = schema_name + + exc_handler = self.handle_database_exceptions() + async with exc_handler, self.with_cursor(self.connection) as cursor: + await cursor.copy_records_to_table(table_name, coerced_records, **copy_kwargs) + if exc_handler.pending_exception is not None: + raise exc_handler.pending_exception from None + + telemetry_payload: StorageTelemetry = { + "destination": table, + "rows_processed": len(records), + "bytes_processed": 0, + } + self._attach_partition_telemetry(telemetry_payload, partitioner) + return self._storage_job(telemetry_payload, telemetry) + async def load_from_arrow( self, table: str, diff --git a/sqlspec/adapters/psqlpy/litestar/store.py b/sqlspec/adapters/psqlpy/litestar/store.py index e4bbaf0f0..0896cc8d6 100644 --- a/sqlspec/adapters/psqlpy/litestar/store.py +++ b/sqlspec/adapters/psqlpy/litestar/store.py @@ -82,31 +82,36 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by Returns: Session data as bytes if found and not expired, None otherwise. """ + if renew_for is not None and self._calculate_expires_at(renew_for) is not None: + new_expires_at = self._calculate_expires_at(renew_for) + sql = f""" + UPDATE {self._table_name} + SET expires_at = $1, updated_at = CURRENT_TIMESTAMP + WHERE session_id = $2 + AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) + RETURNING data + """ + async with self._config.provide_connection() as conn: + single_result = await conn.fetch_row(sql, [new_expires_at, key]) + if not single_result: + return None + row = single_result.result() + if not row: + return None + return bytes(row["data"]) + sql = f""" - SELECT data, expires_at FROM {self._table_name} + SELECT data FROM {self._table_name} WHERE session_id = $1 AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ - async with self._config.provide_connection() as conn: - query_result = await conn.fetch(sql, [key]) - rows = query_result.result() - - if not rows: + single_result = await conn.fetch_row(sql, [key]) + if not single_result: + return None + row = single_result.result() + if not row: return None - - row = rows[0] - - if renew_for is not None and row["expires_at"] is not None: - new_expires_at = self._calculate_expires_at(renew_for) - if new_expires_at is not None: - update_sql = f""" - UPDATE {self._table_name} - SET expires_at = $1, updated_at = CURRENT_TIMESTAMP - WHERE session_id = $2 - """ - await conn.execute(update_sql, [new_expires_at, key]) - return bytes(row["data"]) async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: diff --git a/sqlspec/adapters/psqlpy/type_converter.py b/sqlspec/adapters/psqlpy/type_converter.py index 5f788c0d9..3bda4035c 100644 --- a/sqlspec/adapters/psqlpy/type_converter.py +++ b/sqlspec/adapters/psqlpy/type_converter.py @@ -1,26 +1,36 @@ -"""PostgreSQL-specific helpers for the psqlpy adapter. +"""PostgreSQL-specific helpers for the psqlpy adapter.""" -This module preserves the ``register_pgvector`` placeholder used by the -driver configuration layer. -""" - -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from sqlspec.typing import PGVECTOR_INSTALLED if TYPE_CHECKING: from sqlspec.adapters.psqlpy._typing import PsqlpyConnection as Connection -__all__ = ("register_pgvector",) +__all__ = ("coerce_pgvector", "register_pgvector") + + +def coerce_pgvector(value: Any) -> Any: + """Coerce sequence or numpy array to psqlpy PgVector.""" + if value is None or not PGVECTOR_INSTALLED: + return value + try: + from psqlpy.extra_types import PgVector + + if isinstance(value, PgVector): + return value + if isinstance(value, (list, tuple)): + return PgVector(list(value)) + if hasattr(value, "tolist"): + return PgVector(value.tolist()) + except (ImportError, Exception): + return value + return value def register_pgvector(connection: "Connection") -> None: """Register pgvector type handlers on psqlpy connection. - Currently a placeholder for future implementation. The psqlpy library - does not yet expose a type handler registration API compatible with - pgvector's automatic conversion system. - Args: connection: Psqlpy connection instance. """ diff --git a/sqlspec/adapters/psycopg/_typing.py b/sqlspec/adapters/psycopg/_typing.py index f59eb308c..b30aa323a 100644 --- a/sqlspec/adapters/psycopg/_typing.py +++ b/sqlspec/adapters/psycopg/_typing.py @@ -26,7 +26,9 @@ from psycopg.sql import Identifier as PsycopgIdentifier from psycopg.types.json import Jsonb as PsycopgJsonb from psycopg_pool import AsyncConnectionPool as PsycopgAsyncConnectionPool +from psycopg_pool import AsyncNullConnectionPool as PsycopgAsyncNullConnectionPool from psycopg_pool import ConnectionPool as PsycopgConnectionPool +from psycopg_pool import NullConnectionPool as PsycopgNullConnectionPool from psycopg_pool.abc import AsyncConnectFailedCB as PsycopgAsyncConnectFailedCB from psycopg_pool.abc import AsyncConnectionCB as PsycopgAsyncConnectionCB from psycopg_pool.abc import ConnectFailedCB as PsycopgConnectFailedCB @@ -65,6 +67,7 @@ "PsycopgAsyncConnectionCB", "PsycopgAsyncConnectionPool", "PsycopgAsyncCursor", + "PsycopgAsyncNullConnectionPool", "PsycopgAsyncRawCursor", "PsycopgAsyncRowFactory", "PsycopgAsyncSessionContext", @@ -79,6 +82,7 @@ "PsycopgJsonb", "PsycopgNativeAsyncConnection", "PsycopgNativeAsyncCursor", + "PsycopgNullConnectionPool", "PsycopgPipelineDriver", "PsycopgProgrammingError", "PsycopgRowFactory", diff --git a/sqlspec/adapters/psycopg/adk/store.py b/sqlspec/adapters/psycopg/adk/store.py index 16ae88874..f72f8af74 100644 --- a/sqlspec/adapters/psycopg/adk/store.py +++ b/sqlspec/adapters/psycopg/adk/store.py @@ -8,6 +8,7 @@ from sqlspec.adapters.psycopg._typing import psycopg_dict_row as dict_row from sqlspec.adapters.psycopg._typing import psycopg_errors as errors from sqlspec.adapters.psycopg._typing import psycopg_sql as pg_sql +from sqlspec.adapters.psycopg.core import pipeline_supported from sqlspec.config import ADKConfig from sqlspec.extensions.adk import ( BaseAsyncADKStore, @@ -414,28 +415,50 @@ async def append_event_and_update_state( event_data_value = event_record["event_data"] jsonb_value = Jsonb(event_data_value) if isinstance(event_data_value, dict) else event_data_value - async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: + async with self._config.provide_connection() as conn: try: - await cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - row = await cur.fetchone() + if pipeline_supported() and hasattr(conn, "pipeline"): + async with conn.pipeline(), conn.cursor(row_factory=dict_row) as cur: + await cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + if app_state is not None: + await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) + row = await cur.fetchone() + else: + async with conn.cursor(row_factory=dict_row) as cur: + await cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + row = await cur.fetchone() + if app_state is not None: + await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) if row is None: _raise_missing_session(session_id) - if app_state is not None: - await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) except Exception: await conn.rollback() raise @@ -904,28 +927,50 @@ def append_event_and_update_state( event_data_value = event_record["event_data"] jsonb_value = Jsonb(event_data_value) if isinstance(event_data_value, dict) else event_data_value - with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: + with self._config.provide_connection() as conn: try: - cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - row = cur.fetchone() + if pipeline_supported() and hasattr(conn, "pipeline"): + with conn.pipeline(), conn.cursor(row_factory=dict_row) as cur: + cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + if app_state is not None: + cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) + row = cur.fetchone() + else: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + row = cur.fetchone() + if app_state is not None: + cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) if row is None: _raise_missing_session(session_id) - if app_state is not None: - cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) except Exception: conn.rollback() raise diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index c8395c7e9..d13517ec7 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -16,7 +16,9 @@ PsycopgSyncSessionContext, ) from sqlspec.adapters.psycopg._typing import PsycopgAsyncConnectionPool as AsyncConnectionPool +from sqlspec.adapters.psycopg._typing import PsycopgAsyncNullConnectionPool as AsyncNullConnectionPool from sqlspec.adapters.psycopg._typing import PsycopgConnectionPool as ConnectionPool +from sqlspec.adapters.psycopg._typing import PsycopgNullConnectionPool as NullConnectionPool from sqlspec.adapters.psycopg.core import apply_driver_features, default_statement_config from sqlspec.adapters.psycopg.driver import ( PsycopgAsyncDriver, @@ -43,6 +45,7 @@ from sqlspec.extensions.events import EventRuntimeHints from sqlspec.typing import ALLOYDB_CONNECTOR_INSTALLED from sqlspec.utils.config_tools import normalize_connection_config +from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Awaitable, Callable, Mapping @@ -139,6 +142,7 @@ class PsycopgPoolParams(PsycopgConnectionParams): close_returns: NotRequired[bool] reconnect_failed: NotRequired["ConnectFailedCB | AsyncConnectFailedCB | None"] kwargs: NotRequired["dict[str, Any]"] + null_pool: NotRequired[bool] class PsycopgDriverFeatures(TypedDict): @@ -183,6 +187,7 @@ class PsycopgDriverFeatures(TypedDict): enable_alloydb_iam_auth: Enable AlloyDB IAM database authentication for sync connector connections. Defaults to False. alloydb_ip_type: AlloyDB connector IP type. Defaults to PRIVATE. + null_pool: Enable NullConnectionPool / AsyncNullConnectionPool for serverless / PgBouncer environments. """ enable_pgvector: NotRequired[bool] @@ -197,6 +202,7 @@ class PsycopgDriverFeatures(TypedDict): alloydb_instance_uri: NotRequired[str] enable_alloydb_iam_auth: NotRequired[bool] alloydb_ip_type: NotRequired[str] + null_pool: NotRequired[bool] def build_connection_config(connection_config: "PsycopgPoolParams | Mapping[str, Any] | None") -> dict[str, Any]: @@ -386,6 +392,20 @@ def _setup_alloydb_connector( database=cast("str | None", database), ) + def get_pool_stats(self) -> "dict[str, Any] | None": + """Return connection pool statistics if available.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "get_stats"): + return cast("dict[str, Any]", pool.get_stats()) + return None + + def pop_pool_stats(self) -> "dict[str, Any] | None": + """Return and reset connection pool statistics if available.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "pop_stats"): + return cast("dict[str, Any]", pool.pop_stats()) + return None + def _create_pool(self) -> "ConnectionPool": """Create the actual connection pool.""" all_config = dict(self.connection_config) @@ -422,10 +442,19 @@ def _create_pool(self) -> "ConnectionPool": self._setup_alloydb_connector(all_config, pool_parameters) conninfo = None + is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) + pool_cls = NullConnectionPool if is_null_pool else ConnectionPool + if is_null_pool: + pool_parameters.pop("min_size", None) + pool_parameters.pop("max_size", None) + pool_parameters.pop("max_idle", None) + pool_parameters.pop("max_waiting", None) + pool_parameters.pop("num_workers", None) + if conninfo: - pool = ConnectionPool(conninfo, kwargs=all_config, **pool_parameters) + pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) else: - pool = ConnectionPool("", kwargs=all_config, **pool_parameters) + pool = pool_cls("", kwargs=all_config, **pool_parameters) return pool @@ -434,7 +463,14 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting - # Detect extensions on first connection, update dialect + from psycopg.types.json import set_json_dumps, set_json_loads + + serializer = self.driver_features.get("json_serializer", to_json) + deserializer = self.driver_features.get("json_deserializer", from_json) + with suppress(Exception): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) + if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -707,6 +743,20 @@ def __init__( **kwargs, ) + def get_pool_stats(self) -> "dict[str, Any] | None": + """Return connection pool statistics if available.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "get_stats"): + return cast("dict[str, Any]", pool.get_stats()) + return None + + def pop_pool_stats(self) -> "dict[str, Any] | None": + """Return and reset connection pool statistics if available.""" + pool = self.connection_instance + if pool is not None and hasattr(pool, "pop_stats"): + return cast("dict[str, Any]", pool.pop_stats()) + return None + async def _create_pool(self) -> "AsyncConnectionPool": """Create the actual async connection pool.""" @@ -738,10 +788,20 @@ async def _create_pool(self) -> "AsyncConnectionPool": conninfo = all_config.pop("conninfo", None) kwargs = all_config.pop("kwargs", {}) all_config.update(kwargs) + + is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) + pool_cls = AsyncNullConnectionPool if is_null_pool else AsyncConnectionPool + if is_null_pool: + pool_parameters.pop("min_size", None) + pool_parameters.pop("max_size", None) + pool_parameters.pop("max_idle", None) + pool_parameters.pop("max_waiting", None) + pool_parameters.pop("num_workers", None) + if conninfo: - pool = AsyncConnectionPool(conninfo, kwargs=all_config, **pool_parameters) + pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) else: - pool = AsyncConnectionPool("", kwargs=all_config, **pool_parameters) + pool = pool_cls("", kwargs=all_config, **pool_parameters) if open_pool is True: await pool.open() @@ -753,7 +813,14 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) - # Detect extensions on first connection, update dialect + from psycopg.types.json import set_json_dumps, set_json_loads + + serializer = self.driver_features.get("json_serializer", to_json) + deserializer = self.driver_features.get("json_deserializer", from_json) + with suppress(Exception): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) + if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) diff --git a/sqlspec/adapters/psycopg/core.py b/sqlspec/adapters/psycopg/core.py index 9f4d35290..a677c8e37 100644 --- a/sqlspec/adapters/psycopg/core.py +++ b/sqlspec/adapters/psycopg/core.py @@ -80,11 +80,21 @@ "resolve_runtime_statement_config", ) -TRANSACTION_STATUS_IDLE = 0 -TRANSACTION_STATUS_ACTIVE = 1 -TRANSACTION_STATUS_INTRANS = 2 -TRANSACTION_STATUS_INERROR = 3 -TRANSACTION_STATUS_UNKNOWN = 4 +TRANSACTION_STATUS_IDLE: int = 0 +TRANSACTION_STATUS_ACTIVE: int = 1 +TRANSACTION_STATUS_INTRANS: int = 2 +TRANSACTION_STATUS_INERROR: int = 3 +TRANSACTION_STATUS_UNKNOWN: int = 4 +try: + from psycopg.pq import TransactionStatus + + TRANSACTION_STATUS_IDLE = int(TransactionStatus.IDLE) + TRANSACTION_STATUS_ACTIVE = int(TransactionStatus.ACTIVE) + TRANSACTION_STATUS_INTRANS = int(TransactionStatus.INTRANS) + TRANSACTION_STATUS_INERROR = int(TransactionStatus.INERROR) + TRANSACTION_STATUS_UNKNOWN = int(TransactionStatus.UNKNOWN) +except (ImportError, AttributeError): + pass class PreparedStackOperation(NamedTuple): @@ -118,9 +128,12 @@ def pipeline_supported() -> bool: return False -def build_copy_from_command(table: str, columns: "list[str]") -> "PsycopgComposed": +def build_copy_from_command(table: str, columns: "list[str]", *, binary: bool = False) -> "PsycopgComposed": + """Build a COPY FROM STDIN command with optional binary format.""" table_identifier = _compose_table_identifier(table) column_sql = PsycopgSQL(", ").join([PsycopgIdentifier(column) for column in columns]) + if binary: + return PsycopgSQL("COPY {} ({}) FROM STDIN WITH (FORMAT BINARY)").format(table_identifier, column_sql) return PsycopgSQL("COPY {} ({}) FROM STDIN").format(table_identifier, column_sql) diff --git a/sqlspec/adapters/psycopg/driver.py b/sqlspec/adapters/psycopg/driver.py index 8c6b36f09..dce1a7ed3 100644 --- a/sqlspec/adapters/psycopg/driver.py +++ b/sqlspec/adapters/psycopg/driver.py @@ -286,7 +286,11 @@ def dispatch_execute_many(self, cursor: Any, statement: "SQL") -> "ExecutionResu return self.create_execution_result(cursor, rowcount_override=0, is_many_result=True) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None - cursor.executemany(sql, prepared_parameters) + if pipeline_supported() and hasattr(self.connection, "pipeline") and not self._transaction_active: + with self.connection.pipeline(): + cursor.executemany(sql, prepared_parameters) + else: + cursor.executemany(sql, prepared_parameters) affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) @@ -336,18 +340,18 @@ def dispatch_special_handling(self, cursor: Any, statement: "SQL") -> "SQLResult copy_data = copy_data[0] if is_copy_from_operation(operation_type): - if isinstance(copy_data, (str, bytes)): - data_to_write = copy_data - elif is_readable(copy_data): - data_to_write = copy_data.read() - else: - data_to_write = str(copy_data) - - if isinstance(data_to_write, str): - data_to_write = data_to_write.encode() - with cursor.copy(sql) as copy_ctx: - copy_ctx.write(data_to_write) + if is_readable(copy_data): + chunk_size = 65536 + while chunk := copy_data.read(chunk_size): + if isinstance(chunk, str): + chunk = chunk.encode("utf-8") + copy_ctx.write(chunk) + else: + data_to_write = copy_data if isinstance(copy_data, (str, bytes)) else str(copy_data) + if isinstance(data_to_write, str): + data_to_write = data_to_write.encode("utf-8") + copy_ctx.write(data_to_write) rows_affected = max(cursor.rowcount, 0) @@ -516,22 +520,26 @@ def load_from_arrow( cursor.execute(truncate_sql) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None - columns, records = self._arrow_table_to_rows(arrow_table) - prepared_records = cast( - "list[Any]", - self.prepare_driver_parameters(records, self.statement_config, is_many=True) - if records and self._arrow_rows_need_preparation(arrow_table) - else records, - ) - if records: + if arrow_table.num_rows > 0: + import pyarrow as pa + + columns = list(arrow_table.column_names) copy_sql = build_copy_from_command(table, columns) exc_handler = self.handle_database_exceptions() with ExitStack() as stack: stack.enter_context(exc_handler) cursor = stack.enter_context(self.with_cursor(self.connection)) copy_ctx = stack.enter_context(cursor.copy(copy_sql)) - for record in prepared_records: - copy_ctx.write_row(record) + needs_prep = self._arrow_rows_need_preparation(arrow_table) + for batch in arrow_table.to_batches(): + batch_table = pa.Table.from_batches([batch]) + _, records = self._arrow_table_to_rows(batch_table) + if needs_prep: + records = cast( + "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) + ) + for record in records: + copy_ctx.write_row(record) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None telemetry_payload = self._ingest_telemetry(arrow_table) @@ -803,7 +811,11 @@ async def dispatch_execute_many(self, cursor: Any, statement: "SQL") -> "Executi return self.create_execution_result(cursor, rowcount_override=0, is_many_result=True) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None - await cursor.executemany(sql, prepared_parameters) + if pipeline_supported() and hasattr(self.connection, "pipeline") and not self._transaction_active: + async with self.connection.pipeline(): + await cursor.executemany(sql, prepared_parameters) + else: + await cursor.executemany(sql, prepared_parameters) affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) @@ -853,18 +865,18 @@ async def dispatch_special_handling(self, cursor: Any, statement: "SQL") -> "SQL copy_data = copy_data[0] if is_copy_from_operation(operation_type): - if isinstance(copy_data, (str, bytes)): - data_to_write = copy_data - elif is_readable(copy_data): - data_to_write = copy_data.read() - else: - data_to_write = str(copy_data) - - if isinstance(data_to_write, str): - data_to_write = data_to_write.encode() - async with cursor.copy(sql) as copy_ctx: - await copy_ctx.write(data_to_write) + if is_readable(copy_data): + chunk_size = 65536 + while chunk := copy_data.read(chunk_size): + if isinstance(chunk, str): + chunk = chunk.encode("utf-8") + await copy_ctx.write(chunk) + else: + data_to_write = copy_data if isinstance(copy_data, (str, bytes)) else str(copy_data) + if isinstance(data_to_write, str): + data_to_write = data_to_write.encode("utf-8") + await copy_ctx.write(data_to_write) rows_affected = max(cursor.rowcount, 0) @@ -1038,22 +1050,26 @@ async def load_from_arrow( await cursor.execute(truncate_sql) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None - columns, records = self._arrow_table_to_rows(arrow_table) - prepared_records = cast( - "list[Any]", - self.prepare_driver_parameters(records, self.statement_config, is_many=True) - if records and self._arrow_rows_need_preparation(arrow_table) - else records, - ) - if records: + if arrow_table.num_rows > 0: + import pyarrow as pa + + columns = list(arrow_table.column_names) copy_sql = build_copy_from_command(table, columns) exc_handler = self.handle_database_exceptions() async with AsyncExitStack() as stack: await stack.enter_async_context(exc_handler) cursor = await stack.enter_async_context(self.with_cursor(self.connection)) copy_ctx = await stack.enter_async_context(cursor.copy(copy_sql)) - for record in prepared_records: - await copy_ctx.write_row(record) + needs_prep = self._arrow_rows_need_preparation(arrow_table) + for batch in arrow_table.to_batches(): + batch_table = pa.Table.from_batches([batch]) + _, records = self._arrow_table_to_rows(batch_table) + if needs_prep: + records = cast( + "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) + ) + for record in records: + await copy_ctx.write_row(record) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None telemetry_payload = self._ingest_telemetry(arrow_table) diff --git a/sqlspec/adapters/psycopg/type_converter.py b/sqlspec/adapters/psycopg/type_converter.py index 9b37165e3..e96da46be 100644 --- a/sqlspec/adapters/psycopg/type_converter.py +++ b/sqlspec/adapters/psycopg/type_converter.py @@ -22,13 +22,15 @@ logger = get_logger(__name__) _pgvector_psycopg: Any | None = import_optional("pgvector.psycopg") +_cached_sync_type_infos: dict[str, Any] = {} +_cached_async_type_infos: dict[str, Any] = {} def register_pgvector_sync(connection: "Connection[Any]") -> None: """Register pgvector type handlers on psycopg sync connection. Enables automatic conversion between NumPy arrays and PostgreSQL vector types - using the pgvector-python library. + using the pgvector-python library with cached TypeInfo lookups. Args: connection: Psycopg sync connection. @@ -39,8 +41,30 @@ def register_pgvector_sync(connection: "Connection[Any]") -> None: if pgvector_psycopg is None: return + if _cached_sync_type_infos: + try: + from pgvector.psycopg.bit import register_bit_info + from pgvector.psycopg.halfvec import register_halfvec_info + from pgvector.psycopg.sparsevec import register_sparsevec_info + from pgvector.psycopg.vector import register_vector_info + + if _cached_sync_type_infos.get("vector") is not None: + register_vector_info(connection, _cached_sync_type_infos["vector"]) + if _cached_sync_type_infos.get("bit") is not None: + register_bit_info(connection, _cached_sync_type_infos["bit"]) + if _cached_sync_type_infos.get("halfvec") is not None: + register_halfvec_info(connection, _cached_sync_type_infos["halfvec"]) + if _cached_sync_type_infos.get("sparsevec") is not None: + register_sparsevec_info(connection, _cached_sync_type_infos["sparsevec"]) + except Exception: + _cached_sync_type_infos.clear() + else: + return + try: pgvector_psycopg.register_vector(connection) + for type_name in ("vector", "bit", "halfvec", "sparsevec"): + _cached_sync_type_infos[type_name] = _fetch_type_info_sync(connection, type_name) except (ValueError, TypeError, ProgrammingError) as error: if _is_missing_vector_error(error): return @@ -53,7 +77,7 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: """Register pgvector type handlers on psycopg async connection. Enables automatic conversion between NumPy arrays and PostgreSQL vector types - using the pgvector-python library. + using the pgvector-python library with cached TypeInfo lookups. Args: connection: Psycopg async connection. @@ -64,9 +88,31 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: if pgvector_psycopg is None: return + if _cached_async_type_infos: + try: + from pgvector.psycopg.bit import register_bit_info + from pgvector.psycopg.halfvec import register_halfvec_info + from pgvector.psycopg.sparsevec import register_sparsevec_info + from pgvector.psycopg.vector import register_vector_info + + if _cached_async_type_infos.get("vector") is not None: + register_vector_info(connection, _cached_async_type_infos["vector"]) + if _cached_async_type_infos.get("bit") is not None: + register_bit_info(connection, _cached_async_type_infos["bit"]) + if _cached_async_type_infos.get("halfvec") is not None: + register_halfvec_info(connection, _cached_async_type_infos["halfvec"]) + if _cached_async_type_infos.get("sparsevec") is not None: + register_sparsevec_info(connection, _cached_async_type_infos["sparsevec"]) + except Exception: + _cached_async_type_infos.clear() + else: + return + try: register_vector_async = pgvector_psycopg.register_vector_async await register_vector_async(connection) + for type_name in ("vector", "bit", "halfvec", "sparsevec"): + _cached_async_type_infos[type_name] = await _fetch_type_info_async(connection, type_name) except (ValueError, TypeError, ProgrammingError) as error: if _is_missing_vector_error(error): return @@ -75,6 +121,26 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: logger.exception("Failed to register pgvector for psycopg async") +def _fetch_type_info_sync(connection: "Connection[Any]", type_name: str) -> Any: + """Fetch TypeInfo synchronously, returning None on failure.""" + from psycopg.types import TypeInfo + + try: + return TypeInfo.fetch(connection, type_name) + except Exception: + return None + + +async def _fetch_type_info_async(connection: "AsyncConnection[Any]", type_name: str) -> Any: + """Fetch TypeInfo asynchronously, returning None on failure.""" + from psycopg.types import TypeInfo + + try: + return await TypeInfo.fetch(connection, type_name) + except Exception: + return None + + def _is_missing_vector_error(error: Exception) -> bool: """Check if error indicates missing vector type in database. From b16a8745599b346966d21c9dd15508e3fa1d88a6 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 01:42:21 +0000 Subject: [PATCH 02/10] fix(adapters): resolve postgres and cockroachdb adapter regressions --- sqlspec/adapters/asyncpg/config.py | 3 + sqlspec/adapters/asyncpg/core.py | 12 +- sqlspec/adapters/psqlpy/adk/store.py | 22 +-- sqlspec/adapters/psqlpy/core.py | 61 +----- sqlspec/adapters/psqlpy/driver.py | 181 +----------------- sqlspec/adapters/psqlpy/litestar/store.py | 23 +-- sqlspec/adapters/psycopg/adk/store.py | 121 ++++-------- sqlspec/adapters/psycopg/config.py | 16 +- sqlspec/adapters/psycopg/type_converter.py | 70 +------ .../unit/adapters/test_asyncpg/test_config.py | 7 + 10 files changed, 96 insertions(+), 420 deletions(-) diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index e04eeacbd..a6cde0f13 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -581,6 +581,9 @@ async def create_connection(self) -> "AsyncpgConnection": for key in _POOL_ONLY_CONFIG_KEYS: config.pop(key, None) + if self.driver_features.get("pgbouncer"): + config["statement_cache_size"] = 0 + if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index 62a494216..4963a0154 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -4,7 +4,7 @@ import datetime import re from collections.abc import Sized -from typing import TYPE_CHECKING, Any, Final, NamedTuple, cast +from typing import TYPE_CHECKING, Any, Final, NamedTuple from sqlspec.adapters.asyncpg._typing import asyncpg_module as asyncpg from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile @@ -131,6 +131,10 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str if found_user: config["user"] = user_val + pgbouncer = config.pop("pgbouncer", None) + if pgbouncer: + config.setdefault("statement_cache_size", 0) + return config @@ -459,13 +463,13 @@ async def _start(self) -> None: self._transaction = None raise - async def fetch_chunk(self) -> "list[Any]": + async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() records = await self._driver._run_with_exception_handler(handler, self._cursor.fetch, self._chunk_size) self._driver._check_pending_exception(handler) - if records is None: + if not records: return [] - return cast("list[Any]", records) + return [dict(record) for record in records] async def close(self, error: bool = False) -> None: self._cursor = None diff --git a/sqlspec/adapters/psqlpy/adk/store.py b/sqlspec/adapters/psqlpy/adk/store.py index 41917eea9..d9503c6e9 100644 --- a/sqlspec/adapters/psqlpy/adk/store.py +++ b/sqlspec/adapters/psqlpy/adk/store.py @@ -101,22 +101,20 @@ async def create_session( VALUES ($1, $2, $3, $4, $5, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) RETURNING id, app_name, user_id, state, create_time, update_time """ - single_result = await conn.fetch_row(sql, [session_id, app_name, user_id, owner_id, state]) + result = await conn.fetch(sql, [session_id, app_name, user_id, owner_id, state]) else: sql = f""" INSERT INTO {self._session_table} (id, app_name, user_id, state, create_time, update_time) VALUES ($1, $2, $3, $4, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) RETURNING id, app_name, user_id, state, create_time, update_time """ - single_result = await conn.fetch_row(sql, [session_id, app_name, user_id, state]) + result = await conn.fetch(sql, [session_id, app_name, user_id, state]) - if not single_result: - msg = "Failed to fetch created session" - raise RuntimeError(msg) - row = single_result.result() - if not row: + rows: list[dict[str, Any]] = result.result() if result else [] + if not rows: msg = "Failed to fetch created session" raise RuntimeError(msg) + row = rows[0] return StoredSession( id=row["id"], app_name=row["app_name"], @@ -145,13 +143,13 @@ async def get_session( try: async with self._config.provide_connection() as conn: - single_result = await conn.fetch_row(sql, [app_name, user_id, session_id]) - if not single_result: - return None - row = single_result.result() - if not row: + result = await conn.fetch(sql, [app_name, user_id, session_id]) + rows: list[dict[str, Any]] = result.result() if result else [] + + if not rows: return None + row = rows[0] return StoredSession( id=row["id"], app_name=row["app_name"], diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index 23762920c..904b7ecba 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -77,7 +77,6 @@ "get_parameter_casts", "is_postgres_extension_active", "prepare_parameters_with_casts", - "records_to_arrow_table", "resolve_postgres_extension_state", "resolve_runtime_statement_config", "split_schema_and_table", @@ -252,69 +251,25 @@ def apply_driver_features( return statement_config, features -def collect_rows(query_result: Any | None, as_records: bool = True) -> "tuple[list[Any], list[str]]": +def collect_rows(query_result: Any | None) -> "tuple[list[dict[str, Any]], list[str]]": """Collect psqlpy rows and column names. Args: query_result: Result returned from cursor.fetch(). - as_records: Whether to return Record objects if available. Returns: Tuple of (rows, column_names). """ - if not query_result: - return [], [] - - if as_records and hasattr(query_result, "records"): - records = cast("list[Any]", query_result.records()) - if not records: - return [], [] - first = records[0] - column_names = list(first.keys()) if hasattr(first, "keys") else [] - return records, column_names - - dict_rows = cast("list[dict[str, Any]]", query_result if isinstance(query_result, list) else query_result.result()) + dict_rows: list[dict[str, Any]] = ( + cast("list[dict[str, Any]]", query_result if isinstance(query_result, list) else query_result.result()) + if query_result + else [] + ) if not dict_rows: return [], [] return dict_rows, list(dict_rows[0]) -def records_to_arrow_table(records: list[Any], columns: list[str], schema: Any = None) -> Any: - """Construct a pyarrow Table from records and column names using columnar arrays. - - Args: - records: List of records or row dictionaries. - columns: Column names corresponding to the records. - schema: Optional pyarrow schema. - - Returns: - A pyarrow Table. - """ - import pyarrow as pa - - if not records: - if schema is not None: - return pa.Table.from_batches([], schema=schema) - return pa.Table.from_arrays([pa.array([]) for _ in columns], names=columns) - - first = records[0] - is_dict = isinstance(first, dict) - if schema is not None: - arrays = [ - pa.array( - [r.get(columns[col_idx]) if is_dict else r[col_idx] for r in records], type=schema.field(col_idx).type - ) - for col_idx in range(len(columns)) - ] - return pa.Table.from_arrays(arrays, schema=schema) - - arrays = [ - pa.array([r.get(columns[col_idx]) if is_dict else r[col_idx] for r in records]) - for col_idx in range(len(columns)) - ] - return pa.Table.from_arrays(arrays, names=columns) - - class PsqlpyStreamSource: """Compiled async chunk source streaming dict rows from a psqlpy server-side cursor. @@ -365,14 +320,12 @@ async def _start(self) -> None: await transaction.rollback() raise - async def fetch_chunk(self) -> "list[Any]": + async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() query_result = await self._driver._run_with_exception_handler(handler, self._cursor.fetchmany, self._chunk_size) self._driver._check_pending_exception(handler) if query_result is None: return [] - if hasattr(query_result, "records"): - return cast("list[Any]", query_result.records()) return cast("list[dict[str, Any]]", query_result.result()) async def close(self, error: bool = False) -> None: diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index 850372ffe..7066a9c24 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -4,8 +4,6 @@ and transaction management. """ -import contextlib -from time import perf_counter from typing import TYPE_CHECKING, Any, cast from mypy_extensions import mypyc_attr @@ -26,23 +24,13 @@ format_table_identifier, get_parameter_casts, prepare_parameters_with_casts, - records_to_arrow_table, split_schema_and_table, ) from sqlspec.adapters.psqlpy.data_dictionary import PsqlpyDataDictionary -from sqlspec.core import ( - SQL, - StackResult, - StatementConfig, - create_arrow_result, - get_cache_config, - register_driver_profile, -) -from sqlspec.core.stack import StatementStack +from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler from sqlspec.driver._common import validate_savepoint_name from sqlspec.exceptions import SQLSpecError -from sqlspec.utils.schema import to_value_type from sqlspec.utils.text import normalize_identifier, quote_identifier if TYPE_CHECKING: @@ -368,173 +356,6 @@ def handle_database_exceptions(self) -> "PsqlpyExceptionHandler": """ return PsqlpyExceptionHandler() - async def execute_stack( - self, stack: "StatementStack", *, continue_on_error: bool = False - ) -> "tuple[StackResult, ...]": - """Execute a StatementStack using psqlpy pipelining when available.""" - if not isinstance(stack, StatementStack) or not stack or self.stack_native_disabled or continue_on_error: - return await super().execute_stack(stack, continue_on_error=continue_on_error) - - queries: list[tuple[str, list[Any] | None]] = [] - prepared_operations: list[tuple[Any, Any]] = [] - - for operation in stack.operations: - kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else {} - config = kwargs.pop("statement_config", None) or self.statement_config - sql_statement = self.prepare_statement( - operation.statement, operation.arguments, statement_config=config, kwargs=kwargs - ) - if sql_statement.is_script or sql_statement.is_many: - return await super().execute_stack(stack, continue_on_error=continue_on_error) - sql, params = self._compiled_sql(sql_statement, config) - p_list = list(params) if isinstance(params, (list, tuple)) else None - queries.append((sql, p_list)) - prepared_operations.append((operation, sql_statement)) - - transaction = self.connection.transaction() - needs_commit = False - if not self._connection_in_transaction(): - await transaction.begin() - needs_commit = True - - results: list[StackResult] = [] - try: - query_results = await transaction.pipeline(queries) - if needs_commit: - await transaction.commit() - for (_op, stmt), q_res in zip(prepared_operations, query_results, strict=False): - rows, column_names = collect_rows(q_res) - exec_result = self.create_execution_result( - self.connection, - selected_data=rows, - column_names=column_names, - data_row_count=len(rows), - is_select_result=stmt.returns_rows(), - ) - sql_result = self.build_statement_result(stmt, exec_result) - results.append(StackResult(result=sql_result)) - except Exception as exc: - if needs_commit: - with contextlib.suppress(Exception): - await transaction.rollback() - msg = f"Pipelined stack execution failed: {exc}" - raise SQLSpecError(msg) from exc - - return tuple(results) - - async def select_to_arrow( - self, - statement: Any, - /, - *parameters: Any, - statement_config: "StatementConfig | None" = None, - return_format: str = "table", - native_only: bool = False, - batch_size: int | None = None, - arrow_schema: Any = None, - **kwargs: Any, - ) -> "ArrowResult": - """Execute a query and return results formatted as Apache Arrow.""" - import pyarrow as pa - - config = statement_config or self.statement_config - sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) - sql, prepared_parameters = self._compiled_sql(sql_statement, config) - params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] - - start_time = perf_counter() - query_result: Any = None - exc_handler = self.handle_database_exceptions() - async with exc_handler, self.with_cursor(self.connection) as cursor: - query_result = await cursor.fetch(sql, params) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - execution_time = perf_counter() - start_time - - records = query_result.records() if hasattr(query_result, "records") else query_result.result() - columns = list(records[0].keys()) if records and hasattr(records[0], "keys") else [] - - table = records_to_arrow_table(records, columns, schema=arrow_schema) - - if return_format == "table": - data: Any = table - elif return_format == "batch": - batches = table.to_batches() - data = batches[0] if batches else pa.RecordBatch.from_arrays([], schema=table.schema) - elif return_format == "batches": - data = table.to_batches(max_chunksize=batch_size) if batch_size else table.to_batches() - elif return_format == "reader": - data = table.to_reader(max_chunksize=batch_size) - else: - data = table - - return create_arrow_result( - statement=sql_statement, - data=data, - rows_affected=len(records), - execution_time=execution_time, - metadata={"columns": columns}, - ) - - async def select_one_or_none( - self, - statement: Any, - /, - *parameters: Any, - schema_type: Any = None, - statement_config: "StatementConfig | None" = None, - **kwargs: Any, - ) -> Any: - """Execute a query returning at most one row using fetch_row fast-path.""" - config = statement_config or self.statement_config - sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) - sql, prepared_parameters = self._compiled_sql(sql_statement, config) - params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] - - single_result: Any = None - exc_handler = self.handle_database_exceptions() - async with exc_handler, self.with_cursor(self.connection) as cursor: - single_result = await cursor.fetch_row(sql, params) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - - if single_result is None: - return None - row_dict = single_result.result() if hasattr(single_result, "result") else dict(cast("Any", single_result)) - if not row_dict: - return None - if schema_type is not None: - return self.to_schema(row_dict, schema_type=schema_type) - return row_dict - - async def select_value( - self, - statement: Any, - /, - *parameters: Any, - value_type: Any = None, - statement_config: "StatementConfig | None" = None, - **kwargs: Any, - ) -> Any: - """Execute a query returning a scalar value using fetch_val fast-path.""" - config = statement_config or self.statement_config - sql_statement = self.prepare_statement(statement, parameters, statement_config=config, kwargs=kwargs) - sql, prepared_parameters = self._compiled_sql(sql_statement, config) - params = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) or [] - - val: Any = None - exc_handler = self.handle_database_exceptions() - async with exc_handler, self.with_cursor(self.connection) as cursor: - val = await cursor.fetch_val(sql, params) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - - if val is None: - return None - if value_type is not None: - return to_value_type(val, value_type) - return val - async def select_to_storage( self, statement: "SQL | str", diff --git a/sqlspec/adapters/psqlpy/litestar/store.py b/sqlspec/adapters/psqlpy/litestar/store.py index 0896cc8d6..664307112 100644 --- a/sqlspec/adapters/psqlpy/litestar/store.py +++ b/sqlspec/adapters/psqlpy/litestar/store.py @@ -86,19 +86,18 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by new_expires_at = self._calculate_expires_at(renew_for) sql = f""" UPDATE {self._table_name} - SET expires_at = $1, updated_at = CURRENT_TIMESTAMP + SET expires_at = CASE WHEN expires_at IS NOT NULL THEN $1 ELSE expires_at END, + updated_at = CURRENT_TIMESTAMP WHERE session_id = $2 AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) RETURNING data """ async with self._config.provide_connection() as conn: - single_result = await conn.fetch_row(sql, [new_expires_at, key]) - if not single_result: + query_result = await conn.fetch(sql, [new_expires_at, key]) + rows = query_result.result() if query_result else [] + if not rows: return None - row = single_result.result() - if not row: - return None - return bytes(row["data"]) + return bytes(rows[0]["data"]) sql = f""" SELECT data FROM {self._table_name} @@ -106,13 +105,11 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ async with self._config.provide_connection() as conn: - single_result = await conn.fetch_row(sql, [key]) - if not single_result: - return None - row = single_result.result() - if not row: + query_result = await conn.fetch(sql, [key]) + rows = query_result.result() if query_result else [] + if not rows: return None - return bytes(row["data"]) + return bytes(rows[0]["data"]) async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Store a session value. diff --git a/sqlspec/adapters/psycopg/adk/store.py b/sqlspec/adapters/psycopg/adk/store.py index f72f8af74..16ae88874 100644 --- a/sqlspec/adapters/psycopg/adk/store.py +++ b/sqlspec/adapters/psycopg/adk/store.py @@ -8,7 +8,6 @@ from sqlspec.adapters.psycopg._typing import psycopg_dict_row as dict_row from sqlspec.adapters.psycopg._typing import psycopg_errors as errors from sqlspec.adapters.psycopg._typing import psycopg_sql as pg_sql -from sqlspec.adapters.psycopg.core import pipeline_supported from sqlspec.config import ADKConfig from sqlspec.extensions.adk import ( BaseAsyncADKStore, @@ -415,50 +414,28 @@ async def append_event_and_update_state( event_data_value = event_record["event_data"] jsonb_value = Jsonb(event_data_value) if isinstance(event_data_value, dict) else event_data_value - async with self._config.provide_connection() as conn: + async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: - if pipeline_supported() and hasattr(conn, "pipeline"): - async with conn.pipeline(), conn.cursor(row_factory=dict_row) as cur: - await cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - if app_state is not None: - await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) - row = await cur.fetchone() - else: - async with conn.cursor(row_factory=dict_row) as cur: - await cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - row = await cur.fetchone() - if app_state is not None: - await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) + await cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + await cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + row = await cur.fetchone() if row is None: _raise_missing_session(session_id) + if app_state is not None: + await cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + await cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) except Exception: await conn.rollback() raise @@ -927,50 +904,28 @@ def append_event_and_update_state( event_data_value = event_record["event_data"] jsonb_value = Jsonb(event_data_value) if isinstance(event_data_value, dict) else event_data_value - with self._config.provide_connection() as conn: + with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: - if pipeline_supported() and hasattr(conn, "pipeline"): - with conn.pipeline(), conn.cursor(row_factory=dict_row) as cur: - cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - if app_state is not None: - cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) - row = cur.fetchone() - else: - with conn.cursor(row_factory=dict_row) as cur: - cur.execute( - insert_query, - ( - event_record["id"], - event_record["app_name"], - event_record["user_id"], - event_record["session_id"], - event_record["invocation_id"], - event_record["timestamp"], - jsonb_value, - ), - ) - cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) - row = cur.fetchone() - if app_state is not None: - cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) - if user_state is not None: - cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) + cur.execute( + insert_query, + ( + event_record["id"], + event_record["app_name"], + event_record["user_id"], + event_record["session_id"], + event_record["invocation_id"], + event_record["timestamp"], + jsonb_value, + ), + ) + cur.execute(update_query, (Jsonb(state), app_name, user_id, session_id)) + row = cur.fetchone() if row is None: _raise_missing_session(session_id) + if app_state is not None: + cur.execute(app_upsert_query, (app_name, Jsonb(app_state))) + if user_state is not None: + cur.execute(user_upsert_query, (app_name, user_id, Jsonb(user_state))) except Exception: conn.rollback() raise diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index d13517ec7..341ba778a 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -463,13 +463,15 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting + from psycopg.adapt import AdaptersMap from psycopg.types.json import set_json_dumps, set_json_loads serializer = self.driver_features.get("json_serializer", to_json) deserializer = self.driver_features.get("json_deserializer", from_json) - with suppress(Exception): - set_json_dumps(serializer, conn) - set_json_loads(deserializer, conn) + if isinstance(getattr(conn, "adapters", None), AdaptersMap): + with suppress(Exception): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) if self._pgvector_available is None: detected_extensions: set[str] = set() @@ -813,13 +815,15 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) + from psycopg.adapt import AdaptersMap from psycopg.types.json import set_json_dumps, set_json_loads serializer = self.driver_features.get("json_serializer", to_json) deserializer = self.driver_features.get("json_deserializer", from_json) - with suppress(Exception): - set_json_dumps(serializer, conn) - set_json_loads(deserializer, conn) + if isinstance(getattr(conn, "adapters", None), AdaptersMap): + with suppress(Exception): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) if self._pgvector_available is None: detected_extensions: set[str] = set() diff --git a/sqlspec/adapters/psycopg/type_converter.py b/sqlspec/adapters/psycopg/type_converter.py index e96da46be..9b37165e3 100644 --- a/sqlspec/adapters/psycopg/type_converter.py +++ b/sqlspec/adapters/psycopg/type_converter.py @@ -22,15 +22,13 @@ logger = get_logger(__name__) _pgvector_psycopg: Any | None = import_optional("pgvector.psycopg") -_cached_sync_type_infos: dict[str, Any] = {} -_cached_async_type_infos: dict[str, Any] = {} def register_pgvector_sync(connection: "Connection[Any]") -> None: """Register pgvector type handlers on psycopg sync connection. Enables automatic conversion between NumPy arrays and PostgreSQL vector types - using the pgvector-python library with cached TypeInfo lookups. + using the pgvector-python library. Args: connection: Psycopg sync connection. @@ -41,30 +39,8 @@ def register_pgvector_sync(connection: "Connection[Any]") -> None: if pgvector_psycopg is None: return - if _cached_sync_type_infos: - try: - from pgvector.psycopg.bit import register_bit_info - from pgvector.psycopg.halfvec import register_halfvec_info - from pgvector.psycopg.sparsevec import register_sparsevec_info - from pgvector.psycopg.vector import register_vector_info - - if _cached_sync_type_infos.get("vector") is not None: - register_vector_info(connection, _cached_sync_type_infos["vector"]) - if _cached_sync_type_infos.get("bit") is not None: - register_bit_info(connection, _cached_sync_type_infos["bit"]) - if _cached_sync_type_infos.get("halfvec") is not None: - register_halfvec_info(connection, _cached_sync_type_infos["halfvec"]) - if _cached_sync_type_infos.get("sparsevec") is not None: - register_sparsevec_info(connection, _cached_sync_type_infos["sparsevec"]) - except Exception: - _cached_sync_type_infos.clear() - else: - return - try: pgvector_psycopg.register_vector(connection) - for type_name in ("vector", "bit", "halfvec", "sparsevec"): - _cached_sync_type_infos[type_name] = _fetch_type_info_sync(connection, type_name) except (ValueError, TypeError, ProgrammingError) as error: if _is_missing_vector_error(error): return @@ -77,7 +53,7 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: """Register pgvector type handlers on psycopg async connection. Enables automatic conversion between NumPy arrays and PostgreSQL vector types - using the pgvector-python library with cached TypeInfo lookups. + using the pgvector-python library. Args: connection: Psycopg async connection. @@ -88,31 +64,9 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: if pgvector_psycopg is None: return - if _cached_async_type_infos: - try: - from pgvector.psycopg.bit import register_bit_info - from pgvector.psycopg.halfvec import register_halfvec_info - from pgvector.psycopg.sparsevec import register_sparsevec_info - from pgvector.psycopg.vector import register_vector_info - - if _cached_async_type_infos.get("vector") is not None: - register_vector_info(connection, _cached_async_type_infos["vector"]) - if _cached_async_type_infos.get("bit") is not None: - register_bit_info(connection, _cached_async_type_infos["bit"]) - if _cached_async_type_infos.get("halfvec") is not None: - register_halfvec_info(connection, _cached_async_type_infos["halfvec"]) - if _cached_async_type_infos.get("sparsevec") is not None: - register_sparsevec_info(connection, _cached_async_type_infos["sparsevec"]) - except Exception: - _cached_async_type_infos.clear() - else: - return - try: register_vector_async = pgvector_psycopg.register_vector_async await register_vector_async(connection) - for type_name in ("vector", "bit", "halfvec", "sparsevec"): - _cached_async_type_infos[type_name] = await _fetch_type_info_async(connection, type_name) except (ValueError, TypeError, ProgrammingError) as error: if _is_missing_vector_error(error): return @@ -121,26 +75,6 @@ async def register_pgvector_async(connection: "AsyncConnection[Any]") -> None: logger.exception("Failed to register pgvector for psycopg async") -def _fetch_type_info_sync(connection: "Connection[Any]", type_name: str) -> Any: - """Fetch TypeInfo synchronously, returning None on failure.""" - from psycopg.types import TypeInfo - - try: - return TypeInfo.fetch(connection, type_name) - except Exception: - return None - - -async def _fetch_type_info_async(connection: "AsyncConnection[Any]", type_name: str) -> Any: - """Fetch TypeInfo asynchronously, returning None on failure.""" - from psycopg.types import TypeInfo - - try: - return await TypeInfo.fetch(connection, type_name) - except Exception: - return None - - def _is_missing_vector_error(error: Exception) -> bool: """Check if error indicates missing vector type in database. diff --git a/tests/unit/adapters/test_asyncpg/test_config.py b/tests/unit/adapters/test_asyncpg/test_config.py index 80bdca140..76e4e189b 100644 --- a/tests/unit/adapters/test_asyncpg/test_config.py +++ b/tests/unit/adapters/test_asyncpg/test_config.py @@ -311,3 +311,10 @@ def test_asyncpg_config_normalizes_aliases() -> None: assert "conninfo" not in config.connection_config assert "dbname" not in config.connection_config assert "username" not in config.connection_config + + +def test_asyncpg_pgbouncer_connection_config_sets_statement_cache_size_zero() -> None: + """Enabling pgbouncer in connection_config should pop pgbouncer and disable statement caching.""" + config = AsyncpgConfig(connection_config={"dsn": "postgresql://localhost:5432/test", "pgbouncer": True}) + assert "pgbouncer" not in config.connection_config + assert config.connection_config["statement_cache_size"] == 0 From b92238c5abe6cec0b92273f3fafab0e4cd075916 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 03:11:38 +0000 Subject: [PATCH 03/10] fix(postgres): fix psqlpy connection tracking, record validation, and psycopg feature registry --- sqlspec/adapters/psqlpy/config.py | 32 +++++++--- sqlspec/adapters/psqlpy/driver.py | 64 +++++++++++-------- .../adapters/_shared/_driver_type_system.py | 1 + .../unit/adapters/test_psqlpy/test_config.py | 26 ++++++++ 4 files changed, 88 insertions(+), 35 deletions(-) diff --git a/sqlspec/adapters/psqlpy/config.py b/sqlspec/adapters/psqlpy/config.py index b5b74db3d..9956e229b 100644 --- a/sqlspec/adapters/psqlpy/config.py +++ b/sqlspec/adapters/psqlpy/config.py @@ -1,5 +1,6 @@ """Psqlpy database configuration.""" +import sys from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast from mypy_extensions import mypyc_attr @@ -146,7 +147,12 @@ async def acquire_connection(self) -> "PsqlpyConnection": ctx = pool.acquire() self._ctx = ctx connection = cast("PsqlpyConnection", await ctx.__aenter__()) - await self._config._ensure_connection(connection) # pyright: ignore[reportPrivateUsage] + try: + await self._config._ensure_connection(connection) # pyright: ignore[reportPrivateUsage] + except BaseException: + await ctx.__aexit__(*sys.exc_info()) + self._ctx = None + raise return connection async def release_connection(self, _conn: "PsqlpyConnection", **kwargs: Any) -> None: @@ -170,9 +176,15 @@ async def __aenter__(self) -> PsqlpyConnection: pool = await self._config.create_pool() self._config.connection_instance = pool - self._ctx = pool.acquire() - connection = await self._ctx.__aenter__() - await self._config._ensure_connection(connection) # pyright: ignore[reportPrivateUsage] + ctx = pool.acquire() + self._ctx = ctx + connection = await ctx.__aenter__() + try: + await self._config._ensure_connection(connection) # pyright: ignore[reportPrivateUsage] + except BaseException: + await ctx.__aexit__(*sys.exc_info()) + self._ctx = None + raise return connection # type: ignore[no-any-return] async def __aexit__( @@ -241,6 +253,7 @@ def __init__( self._user_connection_hook: Callable[[PsqlpyConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) + self._initialized_connection_ids: set[int] = set() self._pgvector_available: bool | None = None self._paradedb_available: bool | None = None self._pg_textsearch_available: bool | None = None @@ -283,11 +296,11 @@ async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: ) self._pg_textsearch_available = is_postgres_extension_active(self.driver_features, "pg_textsearch") - if getattr(connection, "_sqlspec_initialized", False): - return - if self._user_connection_hook is not None: - await self._user_connection_hook(connection) - setattr(connection, "_sqlspec_initialized", True) + conn_id = id(connection) + if conn_id not in self._initialized_connection_ids: + if self._user_connection_hook is not None: + await self._user_connection_hook(connection) + self._initialized_connection_ids.add(conn_id) def get_pool_status(self) -> "dict[str, int] | None": """Return connection pool status metrics if pool is active.""" @@ -321,6 +334,7 @@ async def _close_pool(self) -> None: self.connection_instance.close() self.connection_instance = None + self._initialized_connection_ids.clear() async def create_connection(self) -> "PsqlpyConnection": """Create a single async connection (not from pool). diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index 7066a9c24..f882dad05 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -30,7 +30,7 @@ from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler from sqlspec.driver._common import validate_savepoint_name -from sqlspec.exceptions import SQLSpecError +from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError from sqlspec.utils.text import normalize_identifier, quote_identifier if TYPE_CHECKING: @@ -391,6 +391,42 @@ async def load_from_records( ) -> "StorageBridgeJob": """Load Python records into PostgreSQL via psqlpy binary COPY.""" self._require_capability("arrow_import_enabled") + materialized = list(records) + if not materialized: + msg = "load_from_records requires at least one record." + raise ImproperConfigurationError(msg) + + from collections.abc import Mapping as MappingABC + + first_record = materialized[0] + if isinstance(first_record, MappingABC): + resolved_columns = columns if columns is not None else list(first_record.keys()) + expected_keys = set(resolved_columns) + row_tuples: list[tuple[Any, ...]] = [] + for record in materialized: + if not isinstance(record, MappingABC): + msg = "load_from_records mapping records must all be mappings." + raise ImproperConfigurationError(msg) + if set(record.keys()) != expected_keys: + msg = "load_from_records mapping records must all share the same keys." + raise ImproperConfigurationError(msg) + row_tuples.append(tuple(record[col] for col in resolved_columns)) + else: + if columns is None: + msg = "load_from_records requires columns when records are positional sequences." + raise ImproperConfigurationError(msg) + resolved_columns = columns + row_tuples = [] + for record in materialized: + if isinstance(record, MappingABC): + msg = "load_from_records positional records must all have the same shape." + raise ImproperConfigurationError(msg) + row = tuple(record) + if len(row) != len(resolved_columns): + msg = "load_from_records positional records must match the number of columns." + raise ImproperConfigurationError(msg) + row_tuples.append(row) + if overwrite: qualified = format_table_identifier(table) exc_handler = self.handle_database_exceptions() @@ -399,31 +435,7 @@ async def load_from_records( if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None - if not records: - empty_payload: StorageTelemetry = {"destination": table, "rows_processed": 0, "bytes_processed": 0} - self._attach_partition_telemetry(empty_payload, partitioner) - return self._storage_job(empty_payload, telemetry) - schema_name, table_name = split_schema_and_table(table) - first_record = records[0] - from collections.abc import Mapping as MappingABC - - if columns is None: - if isinstance(first_record, MappingABC): - resolved_columns = list(first_record.keys()) - else: - msg = "columns must be provided when records are sequences" - raise SQLSpecError(msg) - else: - resolved_columns = columns - - if isinstance(first_record, MappingABC): - row_tuples = [ - tuple(r.get(col) for col in resolved_columns) for r in cast("Sequence[Mapping[str, Any]]", records) - ] - else: - row_tuples = [tuple(r) for r in cast("Sequence[Sequence[Any]]", records)] - json_columns = await self._resolve_json_columns(schema_name, table_name) coerced_records = coerce_json_columns(row_tuples, resolved_columns, json_columns) @@ -439,7 +451,7 @@ async def load_from_records( telemetry_payload: StorageTelemetry = { "destination": table, - "rows_processed": len(records), + "rows_processed": len(row_tuples), "bytes_processed": 0, } self._attach_partition_telemetry(telemetry_payload, partitioner) diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 0f3480561..66bb8a7ca 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -219,6 +219,7 @@ class SourceEquivalenceCase: "alloydb_instance_uri", "enable_alloydb_iam_auth", "alloydb_ip_type", + "null_pool", ), "pymssql": ("json_serializer", "json_deserializer", "on_connection_create", "enable_events", "events_backend"), "pymysql": ( diff --git a/tests/unit/adapters/test_psqlpy/test_config.py b/tests/unit/adapters/test_psqlpy/test_config.py index 0b77aa70b..f7a8f324c 100644 --- a/tests/unit/adapters/test_psqlpy/test_config.py +++ b/tests/unit/adapters/test_psqlpy/test_config.py @@ -25,6 +25,8 @@ def result(self) -> list[dict[str, str]]: class _ExtensionConnection: + __slots__ = ("_extension_names", "queries") + def __init__(self, extension_names: set[str]) -> None: self._extension_names = extension_names self.queries: list[tuple[str, list[list[str]]]] = [] @@ -164,6 +166,30 @@ async def test_psqlpy_enable_pg_textsearch_detects_extension_and_promotes_dialec assert config.statement_config.dialect == "pg_textsearch" +@pytest.mark.anyio +async def test_psqlpy_ensure_connection_supports_slotted_connections_and_calls_hook_once() -> None: + """_ensure_connection should track connection IDs without setting attributes on slotted connections.""" + hook_calls: list[object] = [] + + async def on_connection_create(conn: PsqlpyConnection) -> None: + hook_calls.append(conn) + + config = PsqlpyConfig( + driver_features={ + "enable_pgvector": False, + "enable_paradedb": False, + "enable_pg_textsearch": False, + "on_connection_create": on_connection_create, + } + ) + connection = _ExtensionConnection(set()) + + await config._ensure_connection(cast("PsqlpyConnection", connection)) # pyright: ignore[reportPrivateUsage] + await config._ensure_connection(cast("PsqlpyConnection", connection)) # pyright: ignore[reportPrivateUsage] + + assert hook_calls == [connection] + + @pytest.mark.anyio async def test_psqlpy_session_context_resolves_callable_statement_config() -> None: """Session context should call statement_config when it's a callable.""" From d82b9266c122b8119641cd2df00b1b9e48d666a6 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 17:11:01 +0000 Subject: [PATCH 04/10] refactor(postgres): remove unused public methods and internal helper export --- sqlspec/adapters/asyncpg/config.py | 18 ---------- sqlspec/adapters/asyncpg/driver.py | 34 +++---------------- .../adapters/cockroach_psycopg/__init__.py | 8 +---- sqlspec/adapters/psqlpy/config.py | 19 ----------- sqlspec/adapters/psqlpy/driver.py | 19 ----------- sqlspec/adapters/psycopg/config.py | 28 --------------- 6 files changed, 6 insertions(+), 120 deletions(-) diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index a6cde0f13..0120c5705 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -458,24 +458,6 @@ def _setup_alloydb_connector(self, config: "dict[str, Any]") -> None: config["connect"] = _AsyncpgAlloydbConnector(self, user, password, database) - def register_type_codec( - self, - typename: str, - *, - schema: str = "public", - encoder: "Callable[..., Any] | None" = None, - decoder: "Callable[..., Any] | None" = None, - format: str = "text", - ) -> None: - """Register a custom type codec to be applied to all connections.""" - self._custom_type_codecs.append({ - "typename": typename, - "schema": schema, - "encoder": encoder, - "decoder": decoder, - "format": format, - }) - async def _create_pool(self) -> "Pool[Record]": """Create the actual async connection pool.""" config = { diff --git a/sqlspec/adapters/asyncpg/driver.py b/sqlspec/adapters/asyncpg/driver.py index eee6b7055..15ea7fe4e 100644 --- a/sqlspec/adapters/asyncpg/driver.py +++ b/sqlspec/adapters/asyncpg/driver.py @@ -169,7 +169,7 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) if statement.operation_type in {"CREATE", "ALTER", "DROP", "TRUNCATE"}: - self.invalidate_prepared_statements() + self._invalidate_prepared_statements() affected_rows = parse_status(result) @@ -217,7 +217,7 @@ async def dispatch_execute_script(self, cursor: "AsyncpgConnection", statement: last_result = result successful_count += 1 - self.invalidate_prepared_statements() + self._invalidate_prepared_statements() return self.create_execution_result( last_result, statement_count=len(statements), successful_statements=successful_count, is_script_result=True @@ -294,14 +294,14 @@ async def rollback(self) -> None: async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" - self.invalidate_prepared_statements() + self._invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET LOCAL search_path TO {quoted_schema}, "$user", public') async def set_migration_non_transactional_schema(self, schema: str) -> None: """Set the PostgreSQL search path for non-transactional migration SQL.""" - self.invalidate_prepared_statements() + self._invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET search_path TO {quoted_schema}, "$user", public') @@ -474,30 +474,6 @@ async def load_from_records( } return self._storage_job(telemetry_payload) - async def copy_from_table( - self, - table: str, - output: Any, - *, - columns: "list[str] | None" = None, - schema_name: "str | None" = None, - format: str = "text", - delimiter: str = "\t", - null: str = "\\N", - ) -> None: - """Export table contents to output stream or file using PostgreSQL COPY TO STDOUT.""" - table_name, resolved_schema, _ = self._copy_target(table) - schema = schema_name or resolved_schema - await self.connection.copy_from_table( - table_name, - output=output, - columns=columns, - schema_name=schema, - format=format, - delimiter=delimiter, - null=null, - ) - async def load_from_storage( self, table: str, @@ -654,7 +630,7 @@ def _connection_in_transaction(self) -> bool: """Check if connection is in transaction.""" return bool(self.connection.is_in_transaction()) - def invalidate_prepared_statements(self) -> None: + def _invalidate_prepared_statements(self) -> None: """Clear cached prepared statements.""" self._prepared_statements.clear() diff --git a/sqlspec/adapters/cockroach_psycopg/__init__.py b/sqlspec/adapters/cockroach_psycopg/__init__.py index a36a39fd9..963f80a06 100644 --- a/sqlspec/adapters/cockroach_psycopg/__init__.py +++ b/sqlspec/adapters/cockroach_psycopg/__init__.py @@ -12,12 +12,7 @@ CockroachPsycopgSyncConfig, build_connection_config, ) -from sqlspec.adapters.cockroach_psycopg.core import ( - CockroachPsycopgRetryConfig, - as_query, - build_statement_config, - driver_profile, -) +from sqlspec.adapters.cockroach_psycopg.core import CockroachPsycopgRetryConfig, build_statement_config, driver_profile from sqlspec.adapters.cockroach_psycopg.driver import ( CockroachPsycopgAsyncDriver, CockroachPsycopgAsyncExceptionHandler, @@ -40,7 +35,6 @@ "CockroachPsycopgSyncExceptionHandler", "CockroachPsycopgSyncSessionContext", "CockroachSyncConnection", - "as_query", "build_connection_config", "build_statement_config", "driver_profile", diff --git a/sqlspec/adapters/psqlpy/config.py b/sqlspec/adapters/psqlpy/config.py index 9956e229b..dd99bc8e1 100644 --- a/sqlspec/adapters/psqlpy/config.py +++ b/sqlspec/adapters/psqlpy/config.py @@ -302,25 +302,6 @@ async def _ensure_connection(self, connection: "PsqlpyConnection") -> None: await self._user_connection_hook(connection) self._initialized_connection_ids.add(conn_id) - def get_pool_status(self) -> "dict[str, int] | None": - """Return connection pool status metrics if pool is active.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "status"): - status = pool.status() - return { - "max_size": status.max_size, - "size": status.size, - "available": status.available, - "waiting": status.waiting, - } - return None - - def resize_pool(self, new_max_size: int) -> None: - """Dynamically resize the active connection pool.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "resize"): - pool.resize(new_max_size) - async def _create_pool(self) -> "ConnectionPool": """Create the actual async connection pool.""" from sqlspec.adapters.psqlpy._typing import PsqlpyConnectionPool as ConnectionPool diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index f882dad05..62028b3d8 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -29,7 +29,6 @@ from sqlspec.adapters.psqlpy.data_dictionary import PsqlpyDataDictionary from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler -from sqlspec.driver._common import validate_savepoint_name from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError from sqlspec.utils.text import normalize_identifier, quote_identifier @@ -249,24 +248,6 @@ async def rollback(self) -> None: finally: self._transaction_active = False - async def savepoint(self, name: str) -> None: - """Create a savepoint within the current transaction.""" - validate_savepoint_name(name) - quoted_name = quote_identifier(name) - await self.connection.execute(f"SAVEPOINT {quoted_name}") - - async def release_savepoint(self, name: str) -> None: - """Release a savepoint within the current transaction.""" - validate_savepoint_name(name) - quoted_name = quote_identifier(name) - await self.connection.execute(f"RELEASE SAVEPOINT {quoted_name}") - - async def rollback_savepoint(self, name: str) -> None: - """Rollback to a savepoint within the current transaction.""" - validate_savepoint_name(name) - quoted_name = quote_identifier(name) - await self.connection.execute(f"ROLLBACK TO SAVEPOINT {quoted_name}") - async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" normalized_schema = normalize_identifier(schema, "postgres") diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index 341ba778a..1d298849a 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -392,20 +392,6 @@ def _setup_alloydb_connector( database=cast("str | None", database), ) - def get_pool_stats(self) -> "dict[str, Any] | None": - """Return connection pool statistics if available.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "get_stats"): - return cast("dict[str, Any]", pool.get_stats()) - return None - - def pop_pool_stats(self) -> "dict[str, Any] | None": - """Return and reset connection pool statistics if available.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "pop_stats"): - return cast("dict[str, Any]", pool.pop_stats()) - return None - def _create_pool(self) -> "ConnectionPool": """Create the actual connection pool.""" all_config = dict(self.connection_config) @@ -745,20 +731,6 @@ def __init__( **kwargs, ) - def get_pool_stats(self) -> "dict[str, Any] | None": - """Return connection pool statistics if available.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "get_stats"): - return cast("dict[str, Any]", pool.get_stats()) - return None - - def pop_pool_stats(self) -> "dict[str, Any] | None": - """Return and reset connection pool statistics if available.""" - pool = self.connection_instance - if pool is not None and hasattr(pool, "pop_stats"): - return cast("dict[str, Any]", pool.pop_stats()) - return None - async def _create_pool(self) -> "AsyncConnectionPool": """Create the actual async connection pool.""" From 5acf52068fbf6510ba97b90c4714ee5cff4b6bbc Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 20:35:31 +0000 Subject: [PATCH 05/10] refactor(postgres): simplify _typing.py by importing non-generic types at top level --- sqlspec/adapters/asyncpg/_typing.py | 5 ++--- sqlspec/adapters/cockroach_asyncpg/_typing.py | 5 ++--- sqlspec/adapters/cockroach_psycopg/_typing.py | 8 +++----- 3 files changed, 7 insertions(+), 11 deletions(-) diff --git a/sqlspec/adapters/asyncpg/_typing.py b/sqlspec/adapters/asyncpg/_typing.py index 1368732ea..52e6aabb3 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,11 @@ AsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] AsyncpgPool: TypeAlias = Pool[Record] - AsyncpgPostgresError: TypeAlias = PostgresError AsyncpgPreparedStatement: TypeAlias = PreparedStatement[Record] if not TYPE_CHECKING: AsyncpgConnection = PoolConnectionProxy AsyncpgPool = Pool - AsyncpgPostgresError = PostgresError AsyncpgPreparedStatement = PreparedStatement diff --git a/sqlspec/adapters/cockroach_asyncpg/_typing.py b/sqlspec/adapters/cockroach_asyncpg/_typing.py index d3362b6c6..4458e079f 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,10 @@ from sqlspec.core import StatementConfig CockroachAsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] - CockroachAsyncpgPostgresError: TypeAlias = PostgresError CockroachAsyncpgPool: TypeAlias = Pool[Record] if not TYPE_CHECKING: 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..b91c94f7a 100644 --- a/sqlspec/adapters/cockroach_psycopg/_typing.py +++ b/sqlspec/adapters/cockroach_psycopg/_typing.py @@ -9,9 +9,9 @@ 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.crdb import AsyncCrdbConnection, CrdbConnection from psycopg.rows import DictRow as PsycopgDictRow from psycopg.rows import dict_row as cockroach_psycopg_dict_row from psycopg.types.json import Jsonb as CockroachPsycopgJsonb @@ -23,8 +23,6 @@ from types import TracebackType from typing import TypeAlias - from psycopg.crdb import AsyncCrdbConnection, CrdbConnection - from sqlspec.adapters.cockroach_psycopg.driver import CockroachPsycopgAsyncDriver, CockroachPsycopgSyncDriver from sqlspec.core import StatementConfig @@ -34,8 +32,8 @@ CockroachAsyncCursor: TypeAlias = AsyncCursor[PsycopgDictRow] if not TYPE_CHECKING: - CockroachSyncConnection = psycopg_crdb.CrdbConnection - CockroachAsyncConnection = psycopg_crdb.AsyncCrdbConnection + CockroachSyncConnection = CrdbConnection + CockroachAsyncConnection = AsyncCrdbConnection CockroachSyncCursor = Cursor CockroachAsyncCursor = AsyncCursor From d5489d20fef05f3df320f405b81b7fce7a5cf09e Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 15:40:22 +0000 Subject: [PATCH 06/10] fix(postgres): revert harmful pipeline/copy overrides, add psqlpy pipeline execute_stack, and clean up typing --- sqlspec/adapters/asyncpg/core.py | 48 +-- sqlspec/adapters/asyncpg/driver.py | 209 ++----------- sqlspec/adapters/cockroach_asyncpg/core.py | 3 +- .../cockroach_psycopg/data_dictionary.py | 4 - sqlspec/adapters/psqlpy/_typing.py | 29 +- sqlspec/adapters/psqlpy/core.py | 58 +--- sqlspec/adapters/psqlpy/data_dictionary.py | 3 - sqlspec/adapters/psqlpy/driver.py | 295 +++++++++++------- sqlspec/adapters/psqlpy/type_converter.py | 15 +- sqlspec/adapters/psycopg/_typing.py | 31 +- sqlspec/adapters/psycopg/config.py | 8 +- sqlspec/adapters/psycopg/core.py | 6 +- sqlspec/adapters/psycopg/data_dictionary.py | 4 - sqlspec/adapters/psycopg/driver.py | 81 ++--- tests/integration/adapters/_shared/_cases.py | 2 - .../adapters/postgres/asyncpg/test_driver.py | 19 -- .../test_contract_statement_stack_parity.py | 2 +- .../test_psqlpy/test_transaction_state.py | 118 ++++++- 18 files changed, 399 insertions(+), 536 deletions(-) diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index 4963a0154..8c5321a2f 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -4,7 +4,7 @@ import datetime import re from collections.abc import Sized -from typing import TYPE_CHECKING, Any, Final, NamedTuple +from typing import TYPE_CHECKING, Any from sqlspec.adapters.asyncpg._typing import asyncpg_module as asyncpg from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile @@ -42,11 +42,10 @@ if TYPE_CHECKING: from collections.abc import Callable, Mapping - from sqlspec.core import SQL, ParameterStyleConfig, StackOperation + from sqlspec.core import ParameterStyleConfig __all__ = ( "AsyncpgStreamSource", - "NormalizedStackOperation", "apply_driver_features", "build_connection_config", "build_postgres_extension_probe_names", @@ -57,7 +56,6 @@ "create_mapped_exception", "default_statement_config", "driver_profile", - "invoke_prepared_statement", "is_postgres_extension_active", "parse_status", "register_json_codecs", @@ -74,17 +72,6 @@ _PGVECTOR_MISSING_LOGGED = False _JSONB_BINARY_VERSION = b"\x01" - -class NormalizedStackOperation(NamedTuple): - """Normalized execution metadata used for prepared stack operations.""" - - operation: "StackOperation" - statement: "SQL" - sql: str - parameters: "tuple[Any, ...] | dict[str, Any] | None" - - -PREPARED_STATEMENT_CACHE_SIZE: Final[int] = 32 _EXCEPTION_MAPPING_DISPATCHER = TypeDispatcher["tuple[str, type[SQLSpecError], str]"]() @@ -170,37 +157,6 @@ def configure_parameter_serializers( return parameter_config.replace(json_serializer=serializer, json_deserializer=effective_deserializer) -async def invoke_prepared_statement( - prepared: Any, parameters: "tuple[Any, ...] | dict[str, Any] | list[Any] | None", *, fetch: bool -) -> Any: - """Invoke an AsyncPG prepared statement with optional parameters. - - Args: - prepared: AsyncPG prepared statement object. - parameters: Prepared parameters payload. - fetch: Whether to fetch rows. - - Returns: - Query result or status message. - """ - if parameters is None: - if fetch: - return await prepared.fetch() - await prepared.fetch() - return prepared.get_statusmsg() - - if isinstance(parameters, dict): - if fetch: - return await prepared.fetch(**parameters) - await prepared.fetch(**parameters) - return prepared.get_statusmsg() - - if fetch: - return await prepared.fetch(*parameters) - await prepared.fetch(*parameters) - return prepared.get_statusmsg() - - def build_statement_config( *, json_serializer: "Callable[[Any], str] | None" = None, json_deserializer: "Callable[[str], Any] | None" = None ) -> "StatementConfig": diff --git a/sqlspec/adapters/asyncpg/driver.py b/sqlspec/adapters/asyncpg/driver.py index 15ea7fe4e..61f081ceb 100644 --- a/sqlspec/adapters/asyncpg/driver.py +++ b/sqlspec/adapters/asyncpg/driver.py @@ -1,48 +1,27 @@ """AsyncPG PostgreSQL driver implementation for async PostgreSQL operations.""" -import re -from collections import OrderedDict from collections.abc import Mapping from contextlib import suppress from io import BytesIO from typing import TYPE_CHECKING, Any, Final, cast -from mypy_extensions import mypyc_attr from sqlglot import exp, parse_one from sqlglot.errors import ParseError from sqlspec.adapters.asyncpg._typing import AsyncpgCursor, AsyncpgPostgresError, AsyncpgSessionContext from sqlspec.adapters.asyncpg.core import ( - PREPARED_STATEMENT_CACHE_SIZE, AsyncpgStreamSource, - NormalizedStackOperation, collect_rows, create_mapped_exception, default_statement_config, driver_profile, - invoke_prepared_statement, parse_status, resolve_many_rowcount, ) from sqlspec.adapters.asyncpg.data_dictionary import AsyncpgDataDictionary -from sqlspec.core import ( - SQL, - StackResult, - StatementStack, - create_sql_result, - get_cache_config, - is_copy_from_operation, - is_copy_operation, - register_driver_profile, -) -from sqlspec.driver import ( - AsyncDriverAdapterBase, - AsyncRowStream, - BaseAsyncExceptionHandler, - StackExecutionObserver, - describe_stack_statement, -) -from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError, StackExecutionError +from sqlspec.core import SQL, get_cache_config, is_copy_from_operation, is_copy_operation, register_driver_profile +from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler +from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError from sqlspec.utils.logging import get_logger from sqlspec.utils.text import normalize_identifier, quote_identifier from sqlspec.utils.type_guards import has_sqlstate @@ -50,7 +29,7 @@ if TYPE_CHECKING: from collections.abc import Sequence - from sqlspec.adapters.asyncpg._typing import AsyncpgConnection, AsyncpgPreparedStatement + from sqlspec.adapters.asyncpg._typing import AsyncpgConnection from sqlspec.core import ArrowResult, SQLResult, StatementConfig from sqlspec.driver import ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry @@ -58,9 +37,6 @@ __all__ = ("AsyncpgCursor", "AsyncpgDriver", "AsyncpgExceptionHandler", "AsyncpgSessionContext") -_COPY_FROM_STDIN_RE: re.Pattern[str] = re.compile( - r'COPY\s+((?:"[^"]+"|\w+)(?:\.(?:"[^"]+"|\w+))?)(?:\s*\([^)]*\))?\s+FROM\s+STDIN', re.IGNORECASE -) _QUALIFIED_TABLE_NAME_PARTS: Final = 2 logger = get_logger("sqlspec.adapters.asyncpg") @@ -88,7 +64,6 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class AsyncpgDriver(AsyncDriverAdapterBase): """AsyncPG PostgreSQL driver for async database operations. @@ -97,7 +72,7 @@ class AsyncpgDriver(AsyncDriverAdapterBase): and caching, and parameter processing with type coercion. """ - __slots__ = ("_data_dictionary", "_prepared_statements", "_transaction") + __slots__ = ("_data_dictionary", "_transaction") dialect = "postgres" def __init__( @@ -113,7 +88,6 @@ def __init__( super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) self._data_dictionary: AsyncpgDataDictionary | None = None - self._prepared_statements: OrderedDict[str, AsyncpgPreparedStatement] = OrderedDict() self._transaction: Any = None async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") -> "ExecutionResult": @@ -168,9 +142,6 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") else: result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) - if statement.operation_type in {"CREATE", "ALTER", "DROP", "TRUNCATE"}: - self._invalidate_prepared_statements() - affected_rows = parse_status(result) return self.create_execution_result(cursor, rowcount_override=affected_rows) @@ -217,8 +188,6 @@ async def dispatch_execute_script(self, cursor: "AsyncpgConnection", statement: last_result = result successful_count += 1 - self._invalidate_prepared_statements() - return self.create_execution_result( last_result, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) @@ -294,14 +263,12 @@ async def rollback(self) -> None: async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" - self._invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET LOCAL search_path TO {quoted_schema}, "$user", public') async def set_migration_non_transactional_schema(self, schema: str) -> None: """Set the PostgreSQL search path for non-transactional migration SQL.""" - self._invalidate_prepared_statements() normalized_schema = normalize_identifier(schema, "postgres") quoted_schema = quote_identifier(normalized_schema) await self.connection.execute(f'SET search_path TO {quoted_schema}, "$user", public') @@ -335,16 +302,6 @@ def handle_database_exceptions(self) -> "AsyncpgExceptionHandler": """Handle database exceptions with PostgreSQL error codes.""" return AsyncpgExceptionHandler() - async def execute_stack( - self, stack: "StatementStack", *, continue_on_error: bool = False - ) -> "tuple[StackResult, ...]": - """Execute a StatementStack using asyncpg's rapid batching.""" - - if not isinstance(stack, StatementStack) or not stack or self.stack_native_disabled: - return await super().execute_stack(stack, continue_on_error=continue_on_error) - - return await self._execute_stack_native(stack, continue_on_error=continue_on_error) - async def select_to_storage( self, statement: "SQL | str", @@ -393,15 +350,11 @@ async def load_from_arrow( msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - import pyarrow as pa - - for batch in arrow_table.to_batches(): - batch_table = pa.Table.from_batches([batch]) - columns, records = self._arrow_table_to_rows(batch_table) - if records: - await self.connection.copy_records_to_table( - table_name, records=records, columns=columns, schema_name=schema_name - ) + columns, records = self._arrow_table_to_rows(arrow_table) + if records: + await self.connection.copy_records_to_table( + table_name, records=records, columns=columns, schema_name=schema_name + ) telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table self._attach_partition_telemetry(telemetry_payload, partitioner) @@ -414,9 +367,8 @@ async def load_from_records( *, columns: "list[str] | None" = None, overwrite: bool = False, - batch_size: int = 1000, ) -> "StorageBridgeJob": - """Load mapping or positional records directly with binary COPY in batches.""" + """Load mapping or positional records directly with binary COPY.""" self._require_capability("arrow_import_enabled") materialized = list(records) if not materialized: @@ -460,12 +412,9 @@ async def load_from_records( msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - for i in range(0, len(copy_rows), batch_size): - chunk = copy_rows[i : i + batch_size] - await self.connection.copy_records_to_table( - table_name, records=chunk, columns=resolved_columns, schema_name=schema_name - ) - + await self.connection.copy_records_to_table( + table_name, records=copy_rows, columns=resolved_columns, schema_name=schema_name + ) telemetry_payload: StorageTelemetry = { "bytes_processed": 0, "destination": table, @@ -538,118 +487,10 @@ def _copy_target(table: str) -> "tuple[str, str | None, str]": quoted_target = f"{quote_identifier(schema_name)}.{quoted_target}" return table_name, schema_name, quoted_target - async def _execute_stack_native( - self, stack: "StatementStack", *, continue_on_error: bool - ) -> "tuple[StackResult, ...]": - results: list[StackResult] = [] - - transaction_cm = None - if not continue_on_error and not self._connection_in_transaction(): - transaction_cm = self.connection.transaction() - - with StackExecutionObserver(self, stack, continue_on_error, native_pipeline=True) as observer: - if transaction_cm is not None: - async with transaction_cm: - await self._run_stack_operations(stack, continue_on_error, observer, results) - else: - await self._run_stack_operations(stack, continue_on_error, observer, results) - - return tuple(results) - - async def _run_stack_operations( - self, - stack: "StatementStack", - continue_on_error: bool, - observer: "StackExecutionObserver", - results: "list[StackResult]", - ) -> None: - """Run operations for statement stack execution. - - Extracted from _execute_stack_native to avoid closure compilation issues. - """ - for index, operation in enumerate(stack.operations): - try: - normalized: NormalizedStackOperation | None = None - if operation.method == "execute": - kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else {} - statement_config = kwargs.pop("statement_config", None) - config = statement_config or self.statement_config - - sql_statement = self.prepare_statement( - operation.statement, operation.arguments, statement_config=config, kwargs=kwargs - ) - if not sql_statement.is_script and not sql_statement.is_many: - sql_text, prepared_parameters = self._compiled_sql(sql_statement, config) - prepared_parameters = cast("tuple[Any, ...] | dict[str, Any] | None", prepared_parameters) - normalized = NormalizedStackOperation( - operation=operation, statement=sql_statement, sql=sql_text, parameters=prepared_parameters - ) - - if normalized is not None: - stack_result = await self._execute_stack_operation_prepared(normalized) - else: - result = await self._execute_stack_operation(operation) - stack_result = StackResult(result=result) - except Exception as exc: - stack_error = StackExecutionError( - index, - describe_stack_statement(operation.statement), - exc, - adapter=type(self).__name__, - mode="continue-on-error" if continue_on_error else "fail-fast", - ) - if continue_on_error: - await self._rollback_failed_stack() - observer.record_operation_error(stack_error) - results.append(StackResult.from_error(stack_error)) - continue - raise stack_error from exc - - results.append(stack_result) - if continue_on_error: - await self._commit_stack_success() - - async def _execute_stack_operation_prepared(self, normalized: "NormalizedStackOperation") -> StackResult: - prepared = await self._get_prepared_statement(normalized.sql) - metadata = {"prepared_statement": True} - - if normalized.statement.returns_rows(): - rows = await invoke_prepared_statement(prepared, normalized.parameters, fetch=True) - data, _ = collect_rows(rows) - sql_result = create_sql_result( - normalized.statement, data=data, rows_affected=len(data), metadata=metadata, row_format="record" - ) - return StackResult.from_sql_result(sql_result) - - status = await invoke_prepared_statement(prepared, normalized.parameters, fetch=False) - rowcount = parse_status(status) - sql_result = create_sql_result(normalized.statement, rows_affected=rowcount, metadata=metadata) - return StackResult.from_sql_result(sql_result) - def _connection_in_transaction(self) -> bool: """Check if connection is in transaction.""" return bool(self.connection.is_in_transaction()) - def _invalidate_prepared_statements(self) -> None: - """Clear cached prepared statements.""" - self._prepared_statements.clear() - - async def _get_prepared_statement(self, sql: str) -> "AsyncpgPreparedStatement": - """Get or prepare a statement with LRU caching.""" - if self.driver_features.get("pgbouncer"): - return cast("AsyncpgPreparedStatement", await self.connection.prepare(sql)) - - cached = self._prepared_statements.get(sql) - if cached is not None: - self._prepared_statements.move_to_end(sql) - return cached - - prepared = cast("AsyncpgPreparedStatement", await self.connection.prepare(sql)) - self._prepared_statements[sql] = prepared - if len(self._prepared_statements) > PREPARED_STATEMENT_CACHE_SIZE: - self._prepared_statements.popitem(last=False) - return prepared - async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: "SQL") -> None: """Handle PostgreSQL COPY operations. @@ -686,9 +527,7 @@ async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: " schema_name: str | None = None if table_name is None: - match = _COPY_FROM_STDIN_RE.search(sql_text) - if match: - table_name = match.group(1) + table_name = _extract_copy_table_name(statement, sql_text) if table_name is None: msg = "COPY FROM STDIN requires a table name or postgres_copy_table execution argument" @@ -714,6 +553,24 @@ async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: " register_driver_profile("asyncpg", driver_profile) +def _extract_copy_table_name(statement: "SQL", sql_text: str) -> "str | None": + expression = statement.expression + if expression is None: + with suppress(ParseError): + expression = parse_one(sql_text, read="postgres") + if not isinstance(expression, exp.Copy): + return None + target = expression.this + if isinstance(target, exp.Schema): + target = target.this + if not isinstance(target, exp.Table) or not target.name: + return None + schema_name = target.db + if schema_name: + return f"{schema_name}.{target.name}" + return target.name + + def _split_copy_table_name(raw_name: str) -> "tuple[str | None, str]": parts = raw_name.split(".", 1) if len(parts) == _QUALIFIED_TABLE_NAME_PARTS: diff --git a/sqlspec/adapters/cockroach_asyncpg/core.py b/sqlspec/adapters/cockroach_asyncpg/core.py index b5ab58983..452f1cd12 100644 --- a/sqlspec/adapters/cockroach_asyncpg/core.py +++ b/sqlspec/adapters/cockroach_asyncpg/core.py @@ -8,6 +8,7 @@ from sqlglot import tokenize from sqlglot.tokenizer_core import TokenType +from sqlspec.adapters.asyncpg.core import build_connection_config as asyncpg_build_connection_config from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate @@ -69,8 +70,6 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachAsyncp def build_connection_config(config: "dict[str, Any]") -> "dict[str, Any]": """Prepare CockroachDB AsyncPG connection config, extracting multi-region server settings.""" - from sqlspec.adapters.asyncpg.core import build_connection_config as asyncpg_build_connection_config - result = asyncpg_build_connection_config(config) server_settings = dict(result.get("server_settings") or {}) if "application_name" in result: diff --git a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py index 74e0ba79f..6f525d864 100644 --- a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py +++ b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py @@ -3,8 +3,6 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar -from mypy_extensions import mypyc_attr - from sqlspec.adapters.psycopg.data_dictionary import PsycopgAsyncDataDictionary, PsycopgSyncDataDictionary from sqlspec.data_dictionary import ( ColumnMetadata, @@ -57,7 +55,6 @@ _COCKROACH_SUPPORTED_DOMAINS = frozenset(_COCKROACH_METADATA_DOMAINS) - {"crdb_internal", "system"} -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class CockroachPsycopgSyncDataDictionary(PsycopgSyncDataDictionary): """CockroachDB sync data dictionary.""" @@ -291,7 +288,6 @@ def get_foreign_keys( ) -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class CockroachPsycopgAsyncDataDictionary(PsycopgAsyncDataDictionary): """CockroachDB async data dictionary.""" diff --git a/sqlspec/adapters/psqlpy/_typing.py b/sqlspec/adapters/psqlpy/_typing.py index 3b04e90b0..31ca76935 100644 --- a/sqlspec/adapters/psqlpy/_typing.py +++ b/sqlspec/adapters/psqlpy/_typing.py @@ -16,33 +16,22 @@ class _PsqlpyUnavailableError(Exception): if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias - from psqlpy import Connection as _PsqlpyConnection + from psqlpy import Connection as PsqlpyConnection from psqlpy import ConnectionPool as PsqlpyConnectionPool - from psqlpy import Listener as _PsqlpyListener - from psqlpy.exceptions import ConnectionExecuteError as _PsqlpyConnectionExecuteError - from psqlpy.exceptions import DatabaseError as _PsqlpyDatabaseError - from psqlpy.exceptions import DataError as _PsqlpyDataError - from psqlpy.exceptions import Error as _PsqlpyError - from psqlpy.exceptions import IntegrityError as _PsqlpyIntegrityError - from psqlpy.exceptions import NotSupportedError as _PsqlpyNotSupportedError - from psqlpy.exceptions import OperationalError as _PsqlpyOperationalError + from psqlpy import Listener as PsqlpyListener + from psqlpy.exceptions import ConnectionExecuteError as PsqlpyConnectionExecuteError + from psqlpy.exceptions import DatabaseError as PsqlpyDatabaseError + from psqlpy.exceptions import DataError as PsqlpyDataError + from psqlpy.exceptions import Error as PsqlpyError + from psqlpy.exceptions import IntegrityError as PsqlpyIntegrityError + from psqlpy.exceptions import NotSupportedError as PsqlpyNotSupportedError + from psqlpy.exceptions import OperationalError as PsqlpyOperationalError from psqlpy.extra_types import JSONB as PSQLPY_JSONB from sqlspec.adapters.psqlpy.driver import PsqlpyDriver from sqlspec.core import StatementConfig - PsqlpyConnection: TypeAlias = _PsqlpyConnection - PsqlpyDataError: TypeAlias = _PsqlpyDataError - PsqlpyDatabaseError: TypeAlias = _PsqlpyDatabaseError - PsqlpyConnectionExecuteError: TypeAlias = _PsqlpyConnectionExecuteError - PsqlpyError: TypeAlias = _PsqlpyError - PsqlpyIntegrityError: TypeAlias = _PsqlpyIntegrityError - PsqlpyListener: TypeAlias = _PsqlpyListener - PsqlpyNotSupportedError: TypeAlias = _PsqlpyNotSupportedError - PsqlpyOperationalError: TypeAlias = _PsqlpyOperationalError - if not TYPE_CHECKING: PsqlpyConnection = import_optional_attr("psqlpy", "Connection") or Any diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index 904b7ecba..1fe07aed4 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -13,6 +13,7 @@ from sqlglot.errors import ParseError from sqlspec.adapters.psqlpy._typing import PsqlpyDataError, PsqlpyIntegrityError, PsqlpyOperationalError +from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector from sqlspec.core import ( DriverParameterProfile, ParameterStyle, @@ -104,34 +105,6 @@ _DML_COUNT_CTE_ALIAS: Final = "_sqlspec_affected" _DML_COUNT_COLUMN: Final = "_sqlspec_rows_affected" _DML_COUNT_QUERY_CACHE_SIZE: Final = 1024 -_PSQLPY_ACCEPTED_POOL_KWARGS: Final[frozenset[str]] = frozenset({ - "dsn", - "username", - "password", - "host", - "hosts", - "port", - "ports", - "db_name", - "target_session_attrs", - "options", - "application_name", - "connect_timeout_sec", - "connect_timeout_nanosec", - "tcp_user_timeout_sec", - "tcp_user_timeout_nanosec", - "keepalives", - "keepalives_idle_sec", - "keepalives_idle_nanosec", - "keepalives_interval_sec", - "keepalives_interval_nanosec", - "keepalives_retries", - "load_balance_hosts", - "max_db_pool_size", - "conn_recycling_method", - "ssl_mode", - "ca_file", -}) logger = get_logger("sqlspec.adapters.psqlpy.core") _NUMERIC_COERCE_TYPES: "tuple[type[Any], ...]" = (float, decimal.Decimal, list, tuple, dict) @@ -200,32 +173,7 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str username = config.pop("username", None) or config.pop("user", None) if username is not None: config["username"] = username - max_size = config.pop("max_size", None) or config.pop("max_db_pool_size", None) - if max_size is not None: - config["max_db_pool_size"] = max_size - timeout = ( - config.pop("connect_timeout_sec", None) or config.pop("connect_timeout", None) or config.pop("timeout", None) - ) - if timeout is not None: - config["connect_timeout_sec"] = int(timeout) - - valid_config: dict[str, Any] = {} - extra_params: dict[str, Any] = {} - for key, value in config.items(): - if key in _PSQLPY_ACCEPTED_POOL_KWARGS: - valid_config[key] = value - else: - extra_params[key] = value - - if extra_params and "dsn" in valid_config: - dsn_val = str(valid_config["dsn"]) - if "?" in dsn_val: - query_suffix = "&" + "&".join(f"{k}={v}" for k, v in extra_params.items()) - valid_config["dsn"] = dsn_val + query_suffix - elif dsn_val.startswith(("postgresql://", "postgres://")): - query_suffix = "?" + "&".join(f"{k}={v}" for k, v in extra_params.items()) - valid_config["dsn"] = dsn_val + query_suffix - return valid_config + return config def apply_driver_features( @@ -662,8 +610,6 @@ def _coerce_parameter_for_cast(value: Any, cast_type: str, serializer: "Callable if upper_cast in _TIMESTAMP_CASTS: return _coerce_timestamp_parameter(value) if upper_cast in _VECTOR_CASTS: - from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector - return coerce_pgvector(value) return value diff --git a/sqlspec/adapters/psqlpy/data_dictionary.py b/sqlspec/adapters/psqlpy/data_dictionary.py index 05baa87d9..a1b231d5d 100644 --- a/sqlspec/adapters/psqlpy/data_dictionary.py +++ b/sqlspec/adapters/psqlpy/data_dictionary.py @@ -3,8 +3,6 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast -from mypy_extensions import mypyc_attr - from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -68,7 +66,6 @@ } -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsqlpyDataDictionary(AsyncDataDictionaryBase): """PostgreSQL-specific async data dictionary via psqlpy.""" diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index 62028b3d8..1526f783a 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -4,10 +4,9 @@ and transaction management. """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, cast -from mypy_extensions import mypyc_attr - from sqlspec.adapters.psqlpy._typing import PsqlpyCursor, PsqlpyDatabaseError, PsqlpyError, PsqlpySessionContext from sqlspec.adapters.psqlpy.core import ( _DML_COUNT_COLUMN, @@ -27,13 +26,28 @@ split_schema_and_table, ) from sqlspec.adapters.psqlpy.data_dictionary import PsqlpyDataDictionary -from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile -from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler -from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError +from sqlspec.core import ( + SQL, + StackResult, + StatementConfig, + StatementStack, + get_cache_config, + is_copy_operation, + register_driver_profile, +) +from sqlspec.driver import ( + AsyncDriverAdapterBase, + AsyncRowStream, + BaseAsyncExceptionHandler, + StackExecutionObserver, + describe_stack_statement, +) +from sqlspec.exceptions import SQLSpecError, StackExecutionError +from sqlspec.utils.logging import get_logger from sqlspec.utils.text import normalize_identifier, quote_identifier if TYPE_CHECKING: - from collections.abc import Mapping, Sequence + from collections.abc import Sequence from sqlspec.adapters.psqlpy._typing import PsqlpyConnection from sqlspec.core import ArrowResult, SQLResult @@ -42,6 +56,8 @@ __all__ = ("PsqlpyCursor", "PsqlpyDriver", "PsqlpyExceptionHandler", "PsqlpySessionContext") +logger = get_logger("sqlspec.adapters.psqlpy") + class PsqlpyExceptionHandler(BaseAsyncExceptionHandler): """Async context manager for handling psqlpy database exceptions. @@ -66,7 +82,6 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsqlpyDriver(AsyncDriverAdapterBase): """PostgreSQL driver implementation using psqlpy. @@ -124,29 +139,10 @@ async def dispatch_execute(self, cursor: "PsqlpyConnection", statement: SQL) -> ) if statement.operation_type in {"INSERT", "UPDATE", "DELETE"}: - if "returning" in sql.lower(): - query_result = await cursor.fetch(sql, params) - dict_rows, column_names = collect_rows(query_result) - rows_affected = len(dict_rows) - return self.create_execution_result( - cursor, - selected_data=dict_rows, - column_names=column_names, - data_row_count=rows_affected, - rowcount_override=rows_affected, - is_select_result=statement.returns_rows(), - ) count_sql = _dml_count_query(sql) if count_sql is not None: count_result = await cursor.fetch(count_sql, params) - count_rows, _ = collect_rows(count_result) - if len(count_rows) != 1 or set(count_rows[0]) != {_DML_COUNT_COLUMN}: - msg = "psqlpy DML row count query returned an invalid result" - raise SQLSpecError(msg) - rows_affected = count_rows[0][_DML_COUNT_COLUMN] - if type(rows_affected) is not int or rows_affected < 0: - msg = "psqlpy DML row count query returned an invalid count" - raise SQLSpecError(msg) + rows_affected = _extract_dml_count(count_result) return self.create_execution_result(cursor, rowcount_override=rows_affected) result = await cursor.execute(sql, params) @@ -178,7 +174,7 @@ async def dispatch_execute_many(self, cursor: "PsqlpyConnection", statement: SQL return self.create_execution_result(cursor, rowcount_override=rows_affected, is_many_result=True) async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: SQL) -> "ExecutionResult": - """Execute SQL script with statement splitting or batch execution. + """Execute SQL script with statement splitting and sequential execution. Args: cursor: Psqlpy connection object @@ -191,17 +187,6 @@ async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: S prepared_parameters = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) statement_config = statement.statement_config - if not prepared_parameters and hasattr(cursor, "execute_batch"): - statements = self.split_script_statements(sql, statement_config, strip_trailing_semicolon=True) - exc_handler = self.handle_database_exceptions() - async with exc_handler: - await cursor.execute_batch(sql) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - return self.create_execution_result( - cursor, statement_count=len(statements), successful_statements=len(statements), is_script_result=True - ) - statements = self.split_script_statements(sql, statement_config, strip_trailing_semicolon=True) successful_count = 0 @@ -337,6 +322,147 @@ def handle_database_exceptions(self) -> "PsqlpyExceptionHandler": """ return PsqlpyExceptionHandler() + async def execute_stack( + self, stack: "StatementStack", *, continue_on_error: bool = False + ) -> "tuple[StackResult, ...]": + """Execute a StatementStack using psqlpy transaction pipeline when supported.""" + if ( + not isinstance(stack, StatementStack) + or not stack + or self.stack_native_disabled + or continue_on_error + or not hasattr(self.connection, "transaction") + ): + return await super().execute_stack(stack, continue_on_error=continue_on_error) + + prepared_ops = self._prepare_pipeline_operations(stack) + if prepared_ops is None: + return await super().execute_stack(stack, continue_on_error=continue_on_error) + + return await self._execute_stack_pipeline(stack, prepared_ops) + + def _prepare_pipeline_operations(self, stack: "StatementStack") -> "list[tuple[SQL, str, list[Any], bool]] | None": + """Prepare stack operations for native psqlpy transaction pipeline execution. + + Returns None when any operation in the stack requires the sequential + fallback path (non-execute methods, per-operation statement_config, + scripts, batch operations, COPY operations, or mapping parameters). + """ + prepared: list[tuple[SQL, str, list[Any], bool]] = [] + for operation in stack.operations: + if operation.method != "execute": + return None + if operation.keyword_arguments and "statement_config" in operation.keyword_arguments: + return None + + kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else None + sql_statement = self.prepare_statement( + operation.statement, operation.arguments, statement_config=self.statement_config, kwargs=kwargs + ) + if sql_statement.is_script or sql_statement.is_many or is_copy_operation(sql_statement.operation_type): + return None + + sql, prepared_parameters = self._compiled_sql(sql_statement, self.statement_config) + if isinstance(prepared_parameters, Mapping): + return None + + params = list(prepared_parameters) if isinstance(prepared_parameters, (list, tuple)) else [] + is_dml_count = False + if not sql_statement.returns_rows() and sql_statement.operation_type in {"INSERT", "UPDATE", "DELETE"}: + try: + count_sql = _dml_count_query(sql) + except SQLSpecError: + return None + if count_sql is not None: + sql = count_sql + is_dml_count = True + + prepared.append((sql_statement, sql, params, is_dml_count)) + return prepared + + async def _execute_stack_pipeline( + self, stack: "StatementStack", prepared_ops: "list[tuple[SQL, str, list[Any], bool]]" + ) -> "tuple[StackResult, ...]": + """Execute prepared stack operations through psqlpy's native Rust pipeline.""" + results: list[StackResult] = [] + started_transaction = False + queries: list[tuple[str, list[Any] | None]] = [(sql, params) for _, sql, params, _ in prepared_ops] + + with StackExecutionObserver(self, stack, continue_on_error=False, native_pipeline=True): + try: + if not self._connection_in_transaction(): + await self.begin() + started_transaction = True + + transaction = self.connection.transaction() + exc_handler = self.handle_database_exceptions() + try: + query_results = await self._run_with_exception_handler( + exc_handler, transaction.pipeline, queries, True + ) + self._check_pending_exception(exc_handler) + except Exception as exc: + stack_error = StackExecutionError( + 0, + describe_stack_statement(stack.operations[0].statement), + exc, + adapter=type(self).__name__, + mode="fail-fast", + native_pipeline=True, + ) + raise stack_error from exc + + assert query_results is not None + for index, ((sql_statement, _, _, is_dml_count), query_result) in enumerate( + zip(prepared_ops, query_results, strict=False) + ): + try: + if sql_statement.returns_rows(): + dict_rows, column_names = collect_rows(query_result) + execution_result = self.create_execution_result( + self.connection, + selected_data=dict_rows, + column_names=column_names, + data_row_count=len(dict_rows), + is_select_result=True, + row_format="dict", + ) + elif is_dml_count: + rows_affected = _extract_dml_count(query_result) + execution_result = self.create_execution_result( + self.connection, rowcount_override=rows_affected + ) + else: + rows_affected = extract_rows_affected(query_result) + execution_result = self.create_execution_result( + self.connection, rowcount_override=rows_affected + ) + except Exception as exc: + stack_error = StackExecutionError( + index, + describe_stack_statement(stack.operations[index].statement), + exc, + adapter=type(self).__name__, + mode="fail-fast", + native_pipeline=True, + ) + raise stack_error from exc + + sql_result = self.build_statement_result(sql_statement, execution_result) + results.append(StackResult.from_sql_result(sql_result)) + + if started_transaction: + await self.commit() + except Exception: + if started_transaction: + try: + await self.rollback() + except Exception as rollback_error: + logger.debug("Rollback after psqlpy pipeline failure failed: %s", rollback_error) + raise + + return tuple(results) + async def select_to_storage( self, statement: "SQL | str", @@ -360,84 +486,6 @@ async def select_to_storage( self._attach_partition_telemetry(telemetry_payload, partitioner) return self._storage_job(telemetry_payload, telemetry) - async def load_from_records( - self, - table: str, - records: "Sequence[Mapping[str, Any]] | Sequence[Sequence[Any]]", - *, - columns: "list[str] | None" = None, - overwrite: bool = False, - partitioner: "dict[str, object] | None" = None, - telemetry: "StorageTelemetry | None" = None, - ) -> "StorageBridgeJob": - """Load Python records into PostgreSQL via psqlpy binary COPY.""" - self._require_capability("arrow_import_enabled") - materialized = list(records) - if not materialized: - msg = "load_from_records requires at least one record." - raise ImproperConfigurationError(msg) - - from collections.abc import Mapping as MappingABC - - first_record = materialized[0] - if isinstance(first_record, MappingABC): - resolved_columns = columns if columns is not None else list(first_record.keys()) - expected_keys = set(resolved_columns) - row_tuples: list[tuple[Any, ...]] = [] - for record in materialized: - if not isinstance(record, MappingABC): - msg = "load_from_records mapping records must all be mappings." - raise ImproperConfigurationError(msg) - if set(record.keys()) != expected_keys: - msg = "load_from_records mapping records must all share the same keys." - raise ImproperConfigurationError(msg) - row_tuples.append(tuple(record[col] for col in resolved_columns)) - else: - if columns is None: - msg = "load_from_records requires columns when records are positional sequences." - raise ImproperConfigurationError(msg) - resolved_columns = columns - row_tuples = [] - for record in materialized: - if isinstance(record, MappingABC): - msg = "load_from_records positional records must all have the same shape." - raise ImproperConfigurationError(msg) - row = tuple(record) - if len(row) != len(resolved_columns): - msg = "load_from_records positional records must match the number of columns." - raise ImproperConfigurationError(msg) - row_tuples.append(row) - - if overwrite: - qualified = format_table_identifier(table) - exc_handler = self.handle_database_exceptions() - async with exc_handler, self.with_cursor(self.connection) as cursor: - await cursor.execute(f"TRUNCATE TABLE {qualified}") - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - - schema_name, table_name = split_schema_and_table(table) - json_columns = await self._resolve_json_columns(schema_name, table_name) - coerced_records = coerce_json_columns(row_tuples, resolved_columns, json_columns) - - copy_kwargs: dict[str, Any] = {"columns": resolved_columns} - if schema_name: - copy_kwargs["schema_name"] = schema_name - - exc_handler = self.handle_database_exceptions() - async with exc_handler, self.with_cursor(self.connection) as cursor: - await cursor.copy_records_to_table(table_name, coerced_records, **copy_kwargs) - if exc_handler.pending_exception is not None: - raise exc_handler.pending_exception from None - - telemetry_payload: StorageTelemetry = { - "destination": table, - "rows_processed": len(row_tuples), - "bytes_processed": 0, - } - self._attach_partition_telemetry(telemetry_payload, partitioner) - return self._storage_job(telemetry_payload, telemetry) - async def load_from_arrow( self, table: str, @@ -580,4 +628,17 @@ def _connection_in_transaction(self) -> bool: return self._transaction_active +def _extract_dml_count(count_result: Any) -> int: + """Validate and extract the affected row count from a psqlpy DML count CTE query result.""" + count_rows, _ = collect_rows(count_result) + if len(count_rows) != 1 or set(count_rows[0]) != {_DML_COUNT_COLUMN}: + msg = "psqlpy DML row count query returned an invalid result" + raise SQLSpecError(msg) + rows_affected = count_rows[0][_DML_COUNT_COLUMN] + if type(rows_affected) is not int or rows_affected < 0: + msg = "psqlpy DML row count query returned an invalid count" + raise SQLSpecError(msg) + return rows_affected + + register_driver_profile("psqlpy", driver_profile) diff --git a/sqlspec/adapters/psqlpy/type_converter.py b/sqlspec/adapters/psqlpy/type_converter.py index 3bda4035c..5078bb1a8 100644 --- a/sqlspec/adapters/psqlpy/type_converter.py +++ b/sqlspec/adapters/psqlpy/type_converter.py @@ -2,7 +2,7 @@ from typing import TYPE_CHECKING, Any -from sqlspec.typing import PGVECTOR_INSTALLED +from sqlspec.typing import PGVECTOR_INSTALLED, import_optional_attr if TYPE_CHECKING: from sqlspec.adapters.psqlpy._typing import PsqlpyConnection as Connection @@ -14,16 +14,17 @@ def coerce_pgvector(value: Any) -> Any: """Coerce sequence or numpy array to psqlpy PgVector.""" if value is None or not PGVECTOR_INSTALLED: return value + pg_vector_cls = import_optional_attr("psqlpy.extra_types", "PgVector") + if pg_vector_cls is None: + return value try: - from psqlpy.extra_types import PgVector - - if isinstance(value, PgVector): + if isinstance(value, pg_vector_cls): return value if isinstance(value, (list, tuple)): - return PgVector(list(value)) + return pg_vector_cls(list(value)) if hasattr(value, "tolist"): - return PgVector(value.tolist()) - except (ImportError, Exception): + return pg_vector_cls(value.tolist()) + except Exception: return value return value diff --git a/sqlspec/adapters/psycopg/_typing.py b/sqlspec/adapters/psycopg/_typing.py index b30aa323a..323c6b299 100644 --- a/sqlspec/adapters/psycopg/_typing.py +++ b/sqlspec/adapters/psycopg/_typing.py @@ -5,7 +5,7 @@ """ import contextlib -from typing import TYPE_CHECKING, Any, Protocol +from typing import TYPE_CHECKING, Any import psycopg as psycopg_module from psycopg import AsyncConnection, AsyncCursor, Connection, Cursor @@ -44,8 +44,7 @@ from google.cloud.alloydb.connector import Connector as PsycopgAlloydbConnector from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver - from sqlspec.builder import QueryBuilder - from sqlspec.core import SQL, Statement, StatementConfig + from sqlspec.core import StatementConfig PsycopgSyncConnection: TypeAlias = Connection[PsycopgDictRow] PsycopgAsyncConnection: TypeAlias = AsyncConnection[PsycopgDictRow] @@ -83,7 +82,6 @@ "PsycopgNativeAsyncConnection", "PsycopgNativeAsyncCursor", "PsycopgNullConnectionPool", - "PsycopgPipelineDriver", "PsycopgProgrammingError", "PsycopgRowFactory", "PsycopgSQL", @@ -139,31 +137,6 @@ async def __aexit__( await self.cursor.close() -class PsycopgPipelineDriver(Protocol): - """Protocol for psycopg pipeline driver methods used in stack execution.""" - - statement_config: "StatementConfig" - - def prepare_statement( - self, - statement: "SQL | Statement | QueryBuilder", - parameters: Any, - *, - statement_config: "StatementConfig | None" = None, - kwargs: "dict[str, Any] | None" = None, - ) -> "SQL": ... - - def prepare_driver_parameters( - self, - parameters: Any, - statement_config: "StatementConfig", - is_many: bool = False, - prepared_statement: Any | None = None, - ) -> Any: ... - - def _compiled_sql(self, statement: "SQL", statement_config: "StatementConfig") -> "tuple[str, Any]": ... - - class PsycopgSyncSessionContext: """Sync context manager for psycopg sessions. diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index 1d298849a..d493b2480 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -4,6 +4,8 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypedDict, cast from mypy_extensions import mypyc_attr +from psycopg.adapt import AdaptersMap +from psycopg.types.json import set_json_dumps, set_json_loads from typing_extensions import NotRequired, Self from sqlspec.adapters.psycopg._typing import ( @@ -449,9 +451,6 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting - from psycopg.adapt import AdaptersMap - from psycopg.types.json import set_json_dumps, set_json_loads - serializer = self.driver_features.get("json_serializer", to_json) deserializer = self.driver_features.get("json_deserializer", from_json) if isinstance(getattr(conn, "adapters", None), AdaptersMap): @@ -787,9 +786,6 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) - from psycopg.adapt import AdaptersMap - from psycopg.types.json import set_json_dumps, set_json_loads - serializer = self.driver_features.get("json_serializer", to_json) deserializer = self.driver_features.get("json_deserializer", from_json) if isinstance(getattr(conn, "adapters", None), AdaptersMap): diff --git a/sqlspec/adapters/psycopg/core.py b/sqlspec/adapters/psycopg/core.py index a677c8e37..13f806b92 100644 --- a/sqlspec/adapters/psycopg/core.py +++ b/sqlspec/adapters/psycopg/core.py @@ -128,12 +128,10 @@ def pipeline_supported() -> bool: return False -def build_copy_from_command(table: str, columns: "list[str]", *, binary: bool = False) -> "PsycopgComposed": - """Build a COPY FROM STDIN command with optional binary format.""" +def build_copy_from_command(table: str, columns: "list[str]") -> "PsycopgComposed": + """Build a COPY FROM STDIN command.""" table_identifier = _compose_table_identifier(table) column_sql = PsycopgSQL(", ").join([PsycopgIdentifier(column) for column in columns]) - if binary: - return PsycopgSQL("COPY {} ({}) FROM STDIN WITH (FORMAT BINARY)").format(table_identifier, column_sql) return PsycopgSQL("COPY {} ({}) FROM STDIN").format(table_identifier, column_sql) diff --git a/sqlspec/adapters/psycopg/data_dictionary.py b/sqlspec/adapters/psycopg/data_dictionary.py index 80c16accc..47522ce58 100644 --- a/sqlspec/adapters/psycopg/data_dictionary.py +++ b/sqlspec/adapters/psycopg/data_dictionary.py @@ -3,8 +3,6 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast -from mypy_extensions import mypyc_attr - from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -68,7 +66,6 @@ } -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsycopgSyncDataDictionary(SyncDataDictionaryBase): """PostgreSQL-specific sync data dictionary.""" @@ -345,7 +342,6 @@ def get_foreign_keys( ) -@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsycopgAsyncDataDictionary(AsyncDataDictionaryBase): """PostgreSQL-specific async data dictionary.""" diff --git a/sqlspec/adapters/psycopg/driver.py b/sqlspec/adapters/psycopg/driver.py index dce1a7ed3..a52aab061 100644 --- a/sqlspec/adapters/psycopg/driver.py +++ b/sqlspec/adapters/psycopg/driver.py @@ -67,12 +67,37 @@ if TYPE_CHECKING: from collections import abc + from typing import Protocol - from sqlspec.adapters.psycopg._typing import PsycopgPipelineDriver - from sqlspec.core import ArrowResult + from sqlspec.builder import QueryBuilder + from sqlspec.core import ArrowResult, Statement from sqlspec.driver import CachedQuery, ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry + class PsycopgPipelineDriver(Protocol): + """Protocol for psycopg pipeline driver methods used in stack execution.""" + + statement_config: "StatementConfig" + + def prepare_statement( + self, + statement: "SQL | Statement | QueryBuilder", + parameters: Any, + *, + statement_config: "StatementConfig | None" = None, + kwargs: "dict[str, Any] | None" = None, + ) -> "SQL": ... + + def prepare_driver_parameters( + self, + parameters: Any, + statement_config: "StatementConfig", + is_many: bool = False, + prepared_statement: Any | None = None, + ) -> Any: ... + + def _compiled_sql(self, statement: "SQL", statement_config: "StatementConfig") -> "tuple[str, Any]": ... + __all__ = ( "PsycopgAsyncCursor", @@ -286,11 +311,7 @@ def dispatch_execute_many(self, cursor: Any, statement: "SQL") -> "ExecutionResu return self.create_execution_result(cursor, rowcount_override=0, is_many_result=True) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None - if pipeline_supported() and hasattr(self.connection, "pipeline") and not self._transaction_active: - with self.connection.pipeline(): - cursor.executemany(sql, prepared_parameters) - else: - cursor.executemany(sql, prepared_parameters) + cursor.executemany(sql, prepared_parameters) affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) @@ -521,25 +542,19 @@ def load_from_arrow( if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None if arrow_table.num_rows > 0: - import pyarrow as pa - - columns = list(arrow_table.column_names) + columns, records = self._arrow_table_to_rows(arrow_table) + if self._arrow_rows_need_preparation(arrow_table): + records = cast( + "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) + ) copy_sql = build_copy_from_command(table, columns) exc_handler = self.handle_database_exceptions() with ExitStack() as stack: stack.enter_context(exc_handler) cursor = stack.enter_context(self.with_cursor(self.connection)) copy_ctx = stack.enter_context(cursor.copy(copy_sql)) - needs_prep = self._arrow_rows_need_preparation(arrow_table) - for batch in arrow_table.to_batches(): - batch_table = pa.Table.from_batches([batch]) - _, records = self._arrow_table_to_rows(batch_table) - if needs_prep: - records = cast( - "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) - ) - for record in records: - copy_ctx.write_row(record) + for record in records: + copy_ctx.write_row(record) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None telemetry_payload = self._ingest_telemetry(arrow_table) @@ -811,11 +826,7 @@ async def dispatch_execute_many(self, cursor: Any, statement: "SQL") -> "Executi return self.create_execution_result(cursor, rowcount_override=0, is_many_result=True) parameter_count = len(prepared_parameters) if isinstance(prepared_parameters, Sized) else None - if pipeline_supported() and hasattr(self.connection, "pipeline") and not self._transaction_active: - async with self.connection.pipeline(): - await cursor.executemany(sql, prepared_parameters) - else: - await cursor.executemany(sql, prepared_parameters) + await cursor.executemany(sql, prepared_parameters) affected_rows = resolve_many_rowcount(cursor, prepared_parameters, fallback_count=parameter_count) return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True) @@ -1051,25 +1062,19 @@ async def load_from_arrow( if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None if arrow_table.num_rows > 0: - import pyarrow as pa - - columns = list(arrow_table.column_names) + columns, records = self._arrow_table_to_rows(arrow_table) + if self._arrow_rows_need_preparation(arrow_table): + records = cast( + "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) + ) copy_sql = build_copy_from_command(table, columns) exc_handler = self.handle_database_exceptions() async with AsyncExitStack() as stack: await stack.enter_async_context(exc_handler) cursor = await stack.enter_async_context(self.with_cursor(self.connection)) copy_ctx = await stack.enter_async_context(cursor.copy(copy_sql)) - needs_prep = self._arrow_rows_need_preparation(arrow_table) - for batch in arrow_table.to_batches(): - batch_table = pa.Table.from_batches([batch]) - _, records = self._arrow_table_to_rows(batch_table) - if needs_prep: - records = cast( - "list[Any]", self.prepare_driver_parameters(records, self.statement_config, is_many=True) - ) - for record in records: - await copy_ctx.write_row(record) + for record in records: + await copy_ctx.write_row(record) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None telemetry_payload = self._ingest_telemetry(arrow_table) diff --git a/tests/integration/adapters/_shared/_cases.py b/tests/integration/adapters/_shared/_cases.py index b1419ca52..b8e269d2c 100644 --- a/tests/integration/adapters/_shared/_cases.py +++ b/tests/integration/adapters/_shared/_cases.py @@ -744,13 +744,11 @@ def _db2_case(mode: Literal["sync", "async"], marks: tuple[Mark | MarkDecorator, supports_connection_hook=True, config_factory_fixture="lifecycle_config_asyncpg", supports_connection_instance=True, - native_stack_parity_mode="standard", extra_assertions=( "explain_modifiers:postgres", "arrow_specifics:postgres", "execute_many_specifics:postgres", "param_codecs:asyncpg", - "statement_stack:native_fallback_parity", "streaming_native:asyncpg", "stream_error_close:pg", ), diff --git a/tests/integration/adapters/postgres/asyncpg/test_driver.py b/tests/integration/adapters/postgres/asyncpg/test_driver.py index 31239e022..b303bbcb6 100644 --- a/tests/integration/adapters/postgres/asyncpg/test_driver.py +++ b/tests/integration/adapters/postgres/asyncpg/test_driver.py @@ -510,25 +510,6 @@ async def test_asyncpg_statement_stack_continue_on_error_inside_transaction(asyn assert [row["id"] for row in persisted.get_data()] == [1, 2] -async def test_asyncpg_statement_stack_marks_prepared(asyncpg_session: "AsyncpgDriver") -> None: - """Prepared statement metadata should be attached to stack results.""" - - await asyncpg_session.execute_script("DELETE FROM test_table_asyncpg") - - stack = ( - StatementStack() - .push_execute("INSERT INTO test_table_asyncpg (id, name, value) VALUES ($1, $2, $3)", (1, "stack-prepared", 50)) - .push_execute("SELECT value FROM test_table_asyncpg WHERE id = $1", (1,)) - ) - - results = await asyncpg_session.execute_stack(stack) - - assert results[0].metadata is not None - assert results[0].metadata.get("prepared_statement") is True - assert results[1].metadata is not None - assert results[1].metadata.get("prepared_statement") is True - - async def test_asyncpg_pool_concurrency(postgres_service: PostgresService) -> None: """Verify that multiple concurrent calls to provide_pool result in a single pool.""" config_params = AsyncpgPoolConfig( diff --git a/tests/unit/adapters/test_contract_statement_stack_parity.py b/tests/unit/adapters/test_contract_statement_stack_parity.py index 9359afd42..f8d21c94e 100644 --- a/tests/unit/adapters/test_contract_statement_stack_parity.py +++ b/tests/unit/adapters/test_contract_statement_stack_parity.py @@ -7,7 +7,7 @@ from tests.integration.adapters._shared.behaviors import STATEMENT_STACK_SCOPE, validate_extra_assertions PARITY_PROOF_KEY = "statement_stack:native_fallback_parity" -OPTED_IN_CASE_IDS = ("psycopg-sync", "asyncpg-async", "psycopg-async", "oracledb-async") +OPTED_IN_CASE_IDS = ("psycopg-sync", "psycopg-async", "oracledb-async") def test_sync_statement_stack_parity_proof_registered() -> None: diff --git a/tests/unit/adapters/test_psqlpy/test_transaction_state.py b/tests/unit/adapters/test_psqlpy/test_transaction_state.py index 9c48ea87f..b92ea9cdd 100644 --- a/tests/unit/adapters/test_psqlpy/test_transaction_state.py +++ b/tests/unit/adapters/test_psqlpy/test_transaction_state.py @@ -7,9 +7,10 @@ from sqlspec.adapters.psqlpy._typing import PsqlpyDatabaseError from sqlspec.adapters.psqlpy.config import PsqlpyConfig -from sqlspec.adapters.psqlpy.core import PsqlpyStreamSource +from sqlspec.adapters.psqlpy.core import _DML_COUNT_COLUMN, PsqlpyStreamSource from sqlspec.adapters.psqlpy.driver import PsqlpyDriver -from sqlspec.exceptions import SQLSpecError +from sqlspec.core import StatementStack +from sqlspec.exceptions import SQLSpecError, StackExecutionError pytestmark = pytest.mark.anyio @@ -205,3 +206,116 @@ async def test_load_from_arrow_decodes_json_text_for_json_columns() -> None: _table_name, records, _kwargs = connection.copy_calls[0] assert records == [(1, {"name": "alpha"}, '{"not": "json"}')] + + +class _PipelineTransaction: + def __init__( + self, + pipeline_calls: "list[tuple[list[tuple[str, list[Any] | None]], bool]]", + results: "list[Any]", + error: "Exception | None" = None, + ) -> None: + self._pipeline_calls = pipeline_calls + self._results = results + self._error = error + + async def pipeline(self, queries: "list[tuple[str, list[Any] | None]]", prepared: bool = True) -> "list[Any]": + self._pipeline_calls.append((queries, prepared)) + if self._error is not None: + raise self._error + return self._results + + +class _PipelineConnection(_FakeConnection): + def __init__(self, results: "list[Any] | None" = None, error: "Exception | None" = None) -> None: + super().__init__() + self.pipeline_calls: list[tuple[list[tuple[str, list[Any] | None]], bool]] = [] + self.fetch_calls: list[tuple[str, Any]] = [] + self.execute_many_calls: list[tuple[str, Any]] = [] + self._results = results or [] + self._error = error + + def transaction(self) -> _PipelineTransaction: + return _PipelineTransaction(self.pipeline_calls, self._results, self._error) + + async def fetch(self, sql: str, parameters: Any = None) -> Any: + self.fetch_calls.append((sql, parameters)) + if _DML_COUNT_COLUMN in sql: + return SimpleNamespace(result=lambda: [{_DML_COUNT_COLUMN: 1}]) + return SimpleNamespace(result=lambda: [{"id": 1}]) + + async def execute_many(self, sql: str, parameters: Any) -> None: + self.execute_many_calls.append((sql, parameters)) + + +async def test_execute_stack_uses_native_transaction_pipeline() -> None: + """Supported execute stacks should run through connection.transaction().pipeline.""" + connection = _PipelineConnection( + results=[ + SimpleNamespace(result=lambda: [{_DML_COUNT_COLUMN: 2}]), + SimpleNamespace(result=lambda: [{"id": 1, "name": "alpha"}]), + ] + ) + driver = PsqlpyDriver(cast("Any", connection)) + stack = ( + StatementStack() + .push_execute("INSERT INTO items (name) VALUES ($1)", "alpha") + .push_execute("SELECT id, name FROM items WHERE name = $1", "alpha") + ) + + results = await driver.execute_stack(stack) + + assert len(results) == 2 + assert results[0].rows_affected == 2 + assert results[1].result is not None + assert results[1].result.get_data() == [{"id": 1, "name": "alpha"}] + assert len(connection.pipeline_calls) == 1 + queries, prepared = connection.pipeline_calls[0] + assert prepared is True + assert len(queries) == 2 + assert _DML_COUNT_COLUMN in queries[0][0] + assert queries[0][1] == ["alpha"] + assert queries[1] == ("SELECT id, name FROM items WHERE name = $1", ["alpha"]) + assert connection.statements == ["BEGIN", "COMMIT"] + + +async def test_execute_stack_falls_back_when_continue_on_error_or_non_execute() -> None: + """Stacks with continue_on_error or non-execute operations must fall back to sequential execution.""" + connection = _PipelineConnection() + driver = PsqlpyDriver(cast("Any", connection)) + + continue_stack = StatementStack().push_execute("SELECT 1") + await driver.execute_stack(continue_stack, continue_on_error=True) + assert connection.pipeline_calls == [] + assert len(connection.fetch_calls) == 1 + + many_stack = StatementStack().push_execute_many("INSERT INTO items (name) VALUES ($1)", [("a",), ("b",)]) + await driver.execute_stack(many_stack) + assert connection.pipeline_calls == [] + assert len(connection.execute_many_calls) == 1 + + +async def test_execute_stack_falls_back_when_native_stack_disabled() -> None: + """Native stack disablement must bypass connection.transaction().pipeline.""" + connection = _PipelineConnection() + driver = PsqlpyDriver(cast("Any", connection), driver_features={"stack_native_disabled": True}) + stack = StatementStack().push_execute("SELECT 1") + + await driver.execute_stack(stack) + + assert connection.pipeline_calls == [] + assert len(connection.fetch_calls) == 1 + + +async def test_execute_stack_native_pipeline_error_rolls_back_and_wraps() -> None: + """Pipeline failures should roll back owned transactions and raise StackExecutionError.""" + connection = _PipelineConnection(error=PsqlpyDatabaseError("unique constraint violation")) + driver = PsqlpyDriver(cast("Any", connection)) + stack = StatementStack().push_execute("INSERT INTO items (name) VALUES ($1)", "dup") + + with pytest.raises(StackExecutionError) as exc_info: + await driver.execute_stack(stack) + + assert exc_info.value.native_pipeline is True + assert connection.statements == ["BEGIN", "ROLLBACK"] + assert driver._connection_in_transaction() is False From 853d9a4778a77810531d4b3caae7bd9ed57e1ad9 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 20:51:48 +0000 Subject: [PATCH 07/10] fix(postgres): preserve native contracts and trim optimization scope --- sqlspec/adapters/asyncpg/_typing.py | 5 +- sqlspec/adapters/asyncpg/config.py | 28 +-- sqlspec/adapters/asyncpg/core.py | 99 +++++---- sqlspec/adapters/asyncpg/driver.py | 203 ++++++++++++----- sqlspec/adapters/cockroach_asyncpg/_typing.py | 5 +- sqlspec/adapters/cockroach_asyncpg/config.py | 10 +- sqlspec/adapters/cockroach_asyncpg/core.py | 37 +--- .../cockroach_asyncpg/data_dictionary.py | 50 +++-- sqlspec/adapters/cockroach_asyncpg/driver.py | 70 ++++-- sqlspec/adapters/cockroach_psycopg/_typing.py | 8 +- .../adapters/cockroach_psycopg/adk/store.py | 93 ++++---- sqlspec/adapters/cockroach_psycopg/config.py | 42 +++- sqlspec/adapters/cockroach_psycopg/core.py | 76 +------ .../cockroach_psycopg/data_dictionary.py | 120 +++++----- sqlspec/adapters/cockroach_psycopg/driver.py | 138 ++++++++---- .../cockroach_psycopg/litestar/store.py | 33 ++- sqlspec/adapters/psqlpy/_typing.py | 29 ++- sqlspec/adapters/psqlpy/core.py | 27 +-- sqlspec/adapters/psqlpy/data_dictionary.py | 3 + sqlspec/adapters/psqlpy/driver.py | 209 ++---------------- sqlspec/adapters/psqlpy/litestar/store.py | 4 +- sqlspec/adapters/psqlpy/type_converter.py | 35 +-- sqlspec/adapters/psycopg/_typing.py | 35 ++- sqlspec/adapters/psycopg/config.py | 51 +---- sqlspec/adapters/psycopg/core.py | 21 +- sqlspec/adapters/psycopg/data_dictionary.py | 4 + sqlspec/adapters/psycopg/driver.py | 29 +-- tests/integration/adapters/_shared/_cases.py | 2 + .../adapters/_shared/_driver_type_system.py | 1 - .../adapters/postgres/asyncpg/test_driver.py | 19 ++ .../unit/adapters/test_asyncpg/test_config.py | 7 - .../test_contract_statement_stack_parity.py | 2 +- .../test_psqlpy/test_transaction_state.py | 118 +--------- 33 files changed, 695 insertions(+), 918 deletions(-) diff --git a/sqlspec/adapters/asyncpg/_typing.py b/sqlspec/adapters/asyncpg/_typing.py index 52e6aabb3..1368732ea 100644 --- a/sqlspec/adapters/asyncpg/_typing.py +++ b/sqlspec/adapters/asyncpg/_typing.py @@ -8,8 +8,7 @@ import asyncpg as asyncpg_module from asyncpg import Connection as AsyncpgRawConnection -from asyncpg import Pool -from asyncpg import PostgresError as AsyncpgPostgresError +from asyncpg import Pool, PostgresError from asyncpg import Record as AsyncpgRecord from asyncpg import connect as asyncpg_connect from asyncpg import create_pool as asyncpg_create_pool @@ -36,11 +35,13 @@ AsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] AsyncpgPool: TypeAlias = Pool[Record] + AsyncpgPostgresError: TypeAlias = PostgresError AsyncpgPreparedStatement: TypeAlias = PreparedStatement[Record] if not TYPE_CHECKING: AsyncpgConnection = PoolConnectionProxy AsyncpgPool = Pool + AsyncpgPostgresError = PostgresError AsyncpgPreparedStatement = PreparedStatement diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index 0120c5705..f92543131 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -101,7 +101,6 @@ class AsyncpgConnectionConfig(TypedDict): connect_timeout: NotRequired[float] command_timeout: NotRequired[float] statement_cache_size: NotRequired[int] - pgbouncer: NotRequired[bool] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] server_settings: NotRequired["dict[str, str]"] @@ -192,9 +191,6 @@ class AsyncpgDriverFeatures(TypedDict): - "notify_queue": Durable queue plus a PostgreSQL notification wakeup hint - "poll_queue": Durable queue discovered by polling Defaults to "notify". - pgbouncer: Enable PgBouncer transaction-pooling compatibility mode. - Disables server-side prepared statement caching (statement_cache_size=0). - type_codecs: Optional list of custom type codec specifications to register. """ json_serializer: NotRequired["Callable[[Any], str]"] @@ -215,8 +211,6 @@ class AsyncpgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["notify", "notify_queue", "poll_queue"]] connection_instance: NotRequired["AsyncpgPool"] on_connection_create: NotRequired["Callable[[AsyncpgConnection], Awaitable[None]]"] - pgbouncer: NotRequired[bool] - type_codecs: NotRequired["list[dict[str, Any]]"] class _AsyncpgCloudSqlConnector: @@ -340,7 +334,6 @@ def __init__( self._user_connection_hook: Callable[[AsyncpgConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) - self._custom_type_codecs: list[dict[str, Any]] = list(features_dict.pop("type_codecs", None) or []) super().__init__( connection_config=build_connection_config(normalize_connection_config(connection_config)), @@ -464,9 +457,6 @@ async def _create_pool(self) -> "Pool[Record]": key: value for key, value in build_connection_config(self.connection_config).items() if value is not None } - if self.connection_config.get("pgbouncer") or self.driver_features.get("pgbouncer"): - config["statement_cache_size"] = 0 - if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): @@ -477,7 +467,7 @@ async def _create_pool(self) -> "Pool[Record]": return await asyncpg_create_pool(**config) async def _init_connection(self, connection: "AsyncpgConnection") -> None: - """Initialize connection with JSON codecs, pgvector support, custom codecs, and user callback. + """Initialize connection with JSON codecs, pgvector support, and user callback. Args: connection: AsyncPG connection to initialize. @@ -489,6 +479,7 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: decoder=self.driver_features.get("json_deserializer", from_json), ) + # Detect extensions on first connection, update dialect if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -509,17 +500,7 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: if self._pgvector_available: await register_pgvector_support(connection) - for codec in self._custom_type_codecs: - codec_kwargs: dict[str, Any] = { - "schema": codec.get("schema", "public"), - "format": codec.get("format", "text"), - } - if codec.get("encoder") is not None: - codec_kwargs["encoder"] = codec["encoder"] - if codec.get("decoder") is not None: - codec_kwargs["decoder"] = codec["decoder"] - await connection.set_type_codec(codec["typename"], **codec_kwargs) - + # Call user-provided callback after internal setup if self._user_connection_hook is not None: await self._user_connection_hook(connection) @@ -563,9 +544,6 @@ async def create_connection(self) -> "AsyncpgConnection": for key in _POOL_ONLY_CONFIG_KEYS: config.pop(key, None) - if self.driver_features.get("pgbouncer"): - config["statement_cache_size"] = 0 - if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index 8c5321a2f..ef7620600 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -4,7 +4,7 @@ import datetime import re from collections.abc import Sized -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Final, NamedTuple from sqlspec.adapters.asyncpg._typing import asyncpg_module as asyncpg from sqlspec.core import DriverParameterProfile, ParameterStyle, StatementConfig, build_statement_config_from_profile @@ -42,10 +42,11 @@ if TYPE_CHECKING: from collections.abc import Callable, Mapping - from sqlspec.core import ParameterStyleConfig + from sqlspec.core import SQL, ParameterStyleConfig, StackOperation __all__ = ( "AsyncpgStreamSource", + "NormalizedStackOperation", "apply_driver_features", "build_connection_config", "build_postgres_extension_probe_names", @@ -56,6 +57,7 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "invoke_prepared_statement", "is_postgres_extension_active", "parse_status", "register_json_codecs", @@ -72,6 +74,17 @@ _PGVECTOR_MISSING_LOGGED = False _JSONB_BINARY_VERSION = b"\x01" + +class NormalizedStackOperation(NamedTuple): + """Normalized execution metadata used for prepared stack operations.""" + + operation: "StackOperation" + statement: "SQL" + sql: str + parameters: "tuple[Any, ...] | dict[str, Any] | None" + + +PREPARED_STATEMENT_CACHE_SIZE: Final[int] = 32 _EXCEPTION_MAPPING_DISPATCHER = TypeDispatcher["tuple[str, type[SQLSpecError], str]"]() @@ -118,10 +131,6 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str if found_user: config["user"] = user_val - pgbouncer = config.pop("pgbouncer", None) - if pgbouncer: - config.setdefault("statement_cache_size", 0) - return config @@ -157,6 +166,37 @@ def configure_parameter_serializers( return parameter_config.replace(json_serializer=serializer, json_deserializer=effective_deserializer) +async def invoke_prepared_statement( + prepared: Any, parameters: "tuple[Any, ...] | dict[str, Any] | list[Any] | None", *, fetch: bool +) -> Any: + """Invoke an AsyncPG prepared statement with optional parameters. + + Args: + prepared: AsyncPG prepared statement object. + parameters: Prepared parameters payload. + fetch: Whether to fetch rows. + + Returns: + Query result or status message. + """ + if parameters is None: + if fetch: + return await prepared.fetch() + await prepared.fetch() + return prepared.get_statusmsg() + + if isinstance(parameters, dict): + if fetch: + return await prepared.fetch(**parameters) + await prepared.fetch(**parameters) + return prepared.get_statusmsg() + + if fetch: + return await prepared.fetch(*parameters) + await prepared.fetch(*parameters) + return prepared.get_statusmsg() + + def build_statement_config( *, json_serializer: "Callable[[Any], str] | None" = None, json_deserializer: "Callable[[str], Any] | None" = None ) -> "StatementConfig": @@ -257,17 +297,10 @@ def parse_status(status: Any) -> int: if not status or not isinstance(status, str): return 0 - stripped = status.strip() - last_space = stripped.rfind(" ") - if last_space != -1: - token = stripped[last_space + 1 :] - if token.isdigit(): - return int(token) - - match = ASYNC_PG_STATUS_REGEX.match(stripped) + match = ASYNC_PG_STATUS_REGEX.match(status.strip()) if match: groups = match.groups() - if len(groups) >= EXPECTED_REGEX_GROUPS and groups[-1]: + if len(groups) >= EXPECTED_REGEX_GROUPS: try: return int(groups[-1]) except (ValueError, IndexError): @@ -423,8 +456,7 @@ async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() records = await self._driver._run_with_exception_handler(handler, self._cursor.fetch, self._chunk_size) self._driver._check_pending_exception(handler) - if not records: - return [] + assert records is not None return [dict(record) for record in records] async def close(self, error: bool = False) -> None: @@ -497,20 +529,12 @@ def _encode_json_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: return str(encoded).encode("utf-8") -def _decode_json_payload(value: Any, decoder: "Callable[..., Any]") -> Any: - """Decode JSON binary or string payload with zero-copy decoding when possible.""" +def _decode_json_payload(value: Any, decoder: "Callable[[str], Any]") -> Any: if isinstance(value, str): return decoder(value) if isinstance(value, memoryview): - raw_bytes = value.tobytes() - elif isinstance(value, (bytes, bytearray)): - raw_bytes = bytes(value) - else: - raw_bytes = bytes(value) - try: - return decoder(raw_bytes) - except (TypeError, UnicodeDecodeError): - return decoder(raw_bytes.decode("utf-8")) + value = value.tobytes() + return decoder(bytes(value).decode("utf-8")) def _encode_jsonb_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: @@ -520,22 +544,15 @@ def _encode_jsonb_payload(value: Any, encoder: "Callable[[Any], str]") -> bytes: return _JSONB_BINARY_VERSION + payload -def _decode_jsonb_payload(value: Any, decoder: "Callable[..., Any]") -> Any: - """Decode JSONB binary or string payload stripping version prefix when present.""" +def _decode_jsonb_payload(value: Any, decoder: "Callable[[str], Any]") -> Any: if isinstance(value, str): return decoder(value) if isinstance(value, memoryview): - raw_bytes = value.tobytes() - elif isinstance(value, (bytes, bytearray)): - raw_bytes = bytes(value) - else: - raw_bytes = bytes(value) - if raw_bytes.startswith(_JSONB_BINARY_VERSION): - raw_bytes = raw_bytes[1:] - try: - return decoder(raw_bytes) - except (TypeError, UnicodeDecodeError): - return decoder(raw_bytes.decode("utf-8")) + value = value.tobytes() + payload = bytes(value) + if payload.startswith(_JSONB_BINARY_VERSION): + payload = payload[1:] + return decoder(payload.decode("utf-8")) def _create_postgres_error( diff --git a/sqlspec/adapters/asyncpg/driver.py b/sqlspec/adapters/asyncpg/driver.py index 61f081ceb..eccab9969 100644 --- a/sqlspec/adapters/asyncpg/driver.py +++ b/sqlspec/adapters/asyncpg/driver.py @@ -1,5 +1,7 @@ """AsyncPG PostgreSQL driver implementation for async PostgreSQL operations.""" +import re +from collections import OrderedDict from collections.abc import Mapping from contextlib import suppress from io import BytesIO @@ -10,18 +12,36 @@ from sqlspec.adapters.asyncpg._typing import AsyncpgCursor, AsyncpgPostgresError, AsyncpgSessionContext from sqlspec.adapters.asyncpg.core import ( + PREPARED_STATEMENT_CACHE_SIZE, AsyncpgStreamSource, + NormalizedStackOperation, collect_rows, create_mapped_exception, default_statement_config, driver_profile, + invoke_prepared_statement, parse_status, resolve_many_rowcount, ) from sqlspec.adapters.asyncpg.data_dictionary import AsyncpgDataDictionary -from sqlspec.core import SQL, get_cache_config, is_copy_from_operation, is_copy_operation, register_driver_profile -from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler -from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError +from sqlspec.core import ( + SQL, + StackResult, + StatementStack, + create_sql_result, + get_cache_config, + is_copy_from_operation, + is_copy_operation, + register_driver_profile, +) +from sqlspec.driver import ( + AsyncDriverAdapterBase, + AsyncRowStream, + BaseAsyncExceptionHandler, + StackExecutionObserver, + describe_stack_statement, +) +from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError, StackExecutionError from sqlspec.utils.logging import get_logger from sqlspec.utils.text import normalize_identifier, quote_identifier from sqlspec.utils.type_guards import has_sqlstate @@ -29,7 +49,7 @@ if TYPE_CHECKING: from collections.abc import Sequence - from sqlspec.adapters.asyncpg._typing import AsyncpgConnection + from sqlspec.adapters.asyncpg._typing import AsyncpgConnection, AsyncpgPreparedStatement from sqlspec.core import ArrowResult, SQLResult, StatementConfig from sqlspec.driver import ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry @@ -37,6 +57,9 @@ __all__ = ("AsyncpgCursor", "AsyncpgDriver", "AsyncpgExceptionHandler", "AsyncpgSessionContext") +_COPY_FROM_STDIN_RE: re.Pattern[str] = re.compile( + r'COPY\s+((?:"[^"]+"|\w+)(?:\.(?:"[^"]+"|\w+))?)(?:\s*\([^)]*\))?\s+FROM\s+STDIN', re.IGNORECASE +) _QUALIFIED_TABLE_NAME_PARTS: Final = 2 logger = get_logger("sqlspec.adapters.asyncpg") @@ -72,7 +95,7 @@ class AsyncpgDriver(AsyncDriverAdapterBase): and caching, and parameter processing with type coercion. """ - __slots__ = ("_data_dictionary", "_transaction") + __slots__ = ("_data_dictionary", "_prepared_statements", "_transaction") dialect = "postgres" def __init__( @@ -88,6 +111,7 @@ def __init__( super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) self._data_dictionary: AsyncpgDataDictionary | None = None + self._prepared_statements: OrderedDict[str, AsyncpgPreparedStatement] = OrderedDict() self._transaction: Any = None async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") -> "ExecutionResult": @@ -104,24 +128,9 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") """ sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) params: tuple[Any, ...] = cast("tuple[Any, ...]", prepared_parameters) if prepared_parameters else () - execution_args = statement.statement_config.execution_args or {} - driver_args = self.statement_config.execution_args or {} - command_timeout = ( - execution_args.get("timeout") - or execution_args.get("command_timeout") - or driver_args.get("timeout") - or driver_args.get("command_timeout") - ) if statement.returns_rows(): - if command_timeout is not None: - records = ( - await cursor.fetch(sql, *params, timeout=command_timeout) - if params - else await cursor.fetch(sql, timeout=command_timeout) - ) - else: - records = await cursor.fetch(sql, *params) if params else await cursor.fetch(sql) + records = await cursor.fetch(sql, *params) if params else await cursor.fetch(sql) data, column_names = collect_rows(records) return self.create_execution_result( @@ -133,14 +142,7 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") row_format="record", ) - if command_timeout is not None: - result = ( - await cursor.execute(sql, *params, timeout=command_timeout) - if params - else await cursor.execute(sql, timeout=command_timeout) - ) - else: - result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) + result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) affected_rows = parse_status(result) @@ -228,7 +230,8 @@ async def begin(self) -> None: await transaction.start() except AsyncpgPostgresError as e: self._release_failed_transaction_claim(transaction) - raise create_mapped_exception(e) from e + msg = f"Failed to begin async transaction: {e}" + raise SQLSpecError(msg) from e self._transaction = transaction def _release_failed_transaction_claim(self, transaction: Any) -> None: @@ -247,7 +250,8 @@ async def commit(self) -> None: else: await self.connection.execute("COMMIT") except AsyncpgPostgresError as e: - raise create_mapped_exception(e) from e + msg = f"Failed to commit async transaction: {e}" + raise SQLSpecError(msg) from e async def rollback(self) -> None: """Rollback the current transaction.""" @@ -259,7 +263,8 @@ async def rollback(self) -> None: else: await self.connection.execute("ROLLBACK") except AsyncpgPostgresError as e: - raise create_mapped_exception(e) from e + msg = f"Failed to rollback async transaction: {e}" + raise SQLSpecError(msg) from e async def set_migration_session_schema(self, schema: str) -> None: """Set the PostgreSQL search path for migration SQL.""" @@ -302,6 +307,16 @@ def handle_database_exceptions(self) -> "AsyncpgExceptionHandler": """Handle database exceptions with PostgreSQL error codes.""" return AsyncpgExceptionHandler() + async def execute_stack( + self, stack: "StatementStack", *, continue_on_error: bool = False + ) -> "tuple[StackResult, ...]": + """Execute a StatementStack using asyncpg's rapid batching.""" + + if not isinstance(stack, StatementStack) or not stack or self.stack_native_disabled: + return await super().execute_stack(stack, continue_on_error=continue_on_error) + + return await self._execute_stack_native(stack, continue_on_error=continue_on_error) + async def select_to_storage( self, statement: "SQL | str", @@ -349,7 +364,6 @@ async def load_from_arrow( except AsyncpgPostgresError as exc: msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - columns, records = self._arrow_table_to_rows(arrow_table) if records: await self.connection.copy_records_to_table( @@ -411,7 +425,6 @@ async def load_from_records( except AsyncpgPostgresError as exc: msg = f"Failed to truncate table '{table}': {exc}" raise SQLSpecError(msg) from exc - await self.connection.copy_records_to_table( table_name, records=copy_rows, columns=resolved_columns, schema_name=schema_name ) @@ -487,10 +500,110 @@ def _copy_target(table: str) -> "tuple[str, str | None, str]": quoted_target = f"{quote_identifier(schema_name)}.{quoted_target}" return table_name, schema_name, quoted_target + async def _execute_stack_native( + self, stack: "StatementStack", *, continue_on_error: bool + ) -> "tuple[StackResult, ...]": + results: list[StackResult] = [] + + transaction_cm = None + if not continue_on_error and not self._connection_in_transaction(): + transaction_cm = self.connection.transaction() + + with StackExecutionObserver(self, stack, continue_on_error, native_pipeline=True) as observer: + if transaction_cm is not None: + async with transaction_cm: + await self._run_stack_operations(stack, continue_on_error, observer, results) + else: + await self._run_stack_operations(stack, continue_on_error, observer, results) + + return tuple(results) + + async def _run_stack_operations( + self, + stack: "StatementStack", + continue_on_error: bool, + observer: "StackExecutionObserver", + results: "list[StackResult]", + ) -> None: + """Run operations for statement stack execution. + + Extracted from _execute_stack_native to avoid closure compilation issues. + """ + for index, operation in enumerate(stack.operations): + try: + normalized: NormalizedStackOperation | None = None + if operation.method == "execute": + kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else {} + statement_config = kwargs.pop("statement_config", None) + config = statement_config or self.statement_config + + sql_statement = self.prepare_statement( + operation.statement, operation.arguments, statement_config=config, kwargs=kwargs + ) + if not sql_statement.is_script and not sql_statement.is_many: + sql_text, prepared_parameters = self._compiled_sql(sql_statement, config) + prepared_parameters = cast("tuple[Any, ...] | dict[str, Any] | None", prepared_parameters) + normalized = NormalizedStackOperation( + operation=operation, statement=sql_statement, sql=sql_text, parameters=prepared_parameters + ) + + if normalized is not None: + stack_result = await self._execute_stack_operation_prepared(normalized) + else: + result = await self._execute_stack_operation(operation) + stack_result = StackResult(result=result) + except Exception as exc: + stack_error = StackExecutionError( + index, + describe_stack_statement(operation.statement), + exc, + adapter=type(self).__name__, + mode="continue-on-error" if continue_on_error else "fail-fast", + ) + if continue_on_error: + await self._rollback_failed_stack() + observer.record_operation_error(stack_error) + results.append(StackResult.from_error(stack_error)) + continue + raise stack_error from exc + + results.append(stack_result) + if continue_on_error: + await self._commit_stack_success() + + async def _execute_stack_operation_prepared(self, normalized: "NormalizedStackOperation") -> StackResult: + prepared = await self._get_prepared_statement(normalized.sql) + metadata = {"prepared_statement": True} + + if normalized.statement.returns_rows(): + rows = await invoke_prepared_statement(prepared, normalized.parameters, fetch=True) + data, _ = collect_rows(rows) + sql_result = create_sql_result( + normalized.statement, data=data, rows_affected=len(data), metadata=metadata, row_format="record" + ) + return StackResult.from_sql_result(sql_result) + + status = await invoke_prepared_statement(prepared, normalized.parameters, fetch=False) + rowcount = parse_status(status) + sql_result = create_sql_result(normalized.statement, rows_affected=rowcount, metadata=metadata) + return StackResult.from_sql_result(sql_result) + def _connection_in_transaction(self) -> bool: """Check if connection is in transaction.""" return bool(self.connection.is_in_transaction()) + async def _get_prepared_statement(self, sql: str) -> "AsyncpgPreparedStatement": + cached = self._prepared_statements.get(sql) + if cached is not None: + self._prepared_statements.move_to_end(sql) + return cached + + prepared = cast("AsyncpgPreparedStatement", await self.connection.prepare(sql)) + self._prepared_statements[sql] = prepared + if len(self._prepared_statements) > PREPARED_STATEMENT_CACHE_SIZE: + self._prepared_statements.popitem(last=False) + return prepared + async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: "SQL") -> None: """Handle PostgreSQL COPY operations. @@ -527,7 +640,9 @@ async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: " schema_name: str | None = None if table_name is None: - table_name = _extract_copy_table_name(statement, sql_text) + match = _COPY_FROM_STDIN_RE.search(sql_text) + if match: + table_name = match.group(1) if table_name is None: msg = "COPY FROM STDIN requires a table name or postgres_copy_table execution argument" @@ -553,24 +668,6 @@ async def _handle_copy_operation(self, cursor: "AsyncpgConnection", statement: " register_driver_profile("asyncpg", driver_profile) -def _extract_copy_table_name(statement: "SQL", sql_text: str) -> "str | None": - expression = statement.expression - if expression is None: - with suppress(ParseError): - expression = parse_one(sql_text, read="postgres") - if not isinstance(expression, exp.Copy): - return None - target = expression.this - if isinstance(target, exp.Schema): - target = target.this - if not isinstance(target, exp.Table) or not target.name: - return None - schema_name = target.db - if schema_name: - return f"{schema_name}.{target.name}" - return target.name - - def _split_copy_table_name(raw_name: str) -> "tuple[str | None, str]": parts = raw_name.split(".", 1) if len(parts) == _QUALIFIED_TABLE_NAME_PARTS: diff --git a/sqlspec/adapters/cockroach_asyncpg/_typing.py b/sqlspec/adapters/cockroach_asyncpg/_typing.py index 4458e079f..d3362b6c6 100644 --- a/sqlspec/adapters/cockroach_asyncpg/_typing.py +++ b/sqlspec/adapters/cockroach_asyncpg/_typing.py @@ -3,8 +3,7 @@ from typing import TYPE_CHECKING, Any import asyncpg as cockroach_asyncpg_module -from asyncpg import Pool -from asyncpg import PostgresError as CockroachAsyncpgPostgresError +from asyncpg import Pool, PostgresError from asyncpg import Record as CockroachAsyncpgRecord from asyncpg import connect as cockroach_asyncpg_connect from asyncpg import create_pool as cockroach_asyncpg_create_pool @@ -21,10 +20,12 @@ from sqlspec.core import StatementConfig CockroachAsyncpgConnection: TypeAlias = Connection[Record] | PoolConnectionProxy[Record] + CockroachAsyncpgPostgresError: TypeAlias = PostgresError CockroachAsyncpgPool: TypeAlias = Pool[Record] if not TYPE_CHECKING: CockroachAsyncpgConnection = PoolConnectionProxy + CockroachAsyncpgPostgresError = PostgresError CockroachAsyncpgPool = Pool __all__ = ( diff --git a/sqlspec/adapters/cockroach_asyncpg/config.py b/sqlspec/adapters/cockroach_asyncpg/config.py index 72fcaf9d4..158eb7942 100644 --- a/sqlspec/adapters/cockroach_asyncpg/config.py +++ b/sqlspec/adapters/cockroach_asyncpg/config.py @@ -7,6 +7,7 @@ from sqlspec.adapters.asyncpg.core import ( apply_driver_features, + build_connection_config, default_statement_config, register_json_codecs, register_pgvector_support, @@ -19,7 +20,7 @@ from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgRecord as Record from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_connect as asyncpg_connect from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_create_pool as asyncpg_create_pool -from sqlspec.adapters.cockroach_asyncpg.core import build_connection_config, validate_follower_read_staleness +from sqlspec.adapters.cockroach_asyncpg.core import validate_follower_read_staleness from sqlspec.adapters.cockroach_asyncpg.driver import CockroachAsyncpgDriver, CockroachAsyncpgExceptionHandler from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs from sqlspec.core.capabilities import TypeCoercionCapabilities @@ -79,10 +80,6 @@ class CockroachAsyncpgConnectionConfig(TypedDict): timeout: NotRequired[float] connect_timeout: NotRequired[float] command_timeout: NotRequired[float] - application_name: NotRequired[str] - gateway_region: NotRequired[str] - default_transaction_use_follower_reads: NotRequired[bool] - results_buffer_size: NotRequired[int] statement_cache_size: NotRequired[int] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] @@ -178,6 +175,7 @@ class CockroachAsyncpgDriverFeatures(TypedDict): class _CockroachAsyncpgSessionFactory(AsyncPoolSessionFactory): """Uses pool.acquire() context manager pattern instead of direct acquire/release.""" + # _connection inherited from AsyncPoolSessionFactory.__slots__ is never written; this class uses _ctx exclusively via the pool.acquire() context manager pattern. __slots__ = ("_ctx",) def __init__(self, config: "CockroachAsyncpgConfig") -> None: @@ -211,7 +209,7 @@ class CockroachAsyncpgConfig( """Configuration for CockroachDB using AsyncPG.""" driver_type: "ClassVar[type[CockroachAsyncpgDriver]]" = CockroachAsyncpgDriver - connection_type: "ClassVar[type[CockroachAsyncpgConnection]]" = cast("Any", CockroachAsyncpgConnection) + connection_type: "ClassVar[type[CockroachAsyncpgConnection]]" = CockroachAsyncpgConnection # type: ignore[assignment] supports_transactional_ddl: "ClassVar[bool]" = False supports_migration_schemas: "ClassVar[bool]" = True supports_native_arrow_export: "ClassVar[bool]" = True diff --git a/sqlspec/adapters/cockroach_asyncpg/core.py b/sqlspec/adapters/cockroach_asyncpg/core.py index 452f1cd12..fefa80e03 100644 --- a/sqlspec/adapters/cockroach_asyncpg/core.py +++ b/sqlspec/adapters/cockroach_asyncpg/core.py @@ -1,14 +1,13 @@ """CockroachDB AsyncPG adapter helpers.""" +import random import re -import secrets from typing import TYPE_CHECKING, Any, Final, cast from mypy_extensions import mypyc_attr from sqlglot import tokenize from sqlglot.tokenizer_core import TokenType -from sqlspec.adapters.asyncpg.core import build_connection_config as asyncpg_build_connection_config from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate @@ -20,7 +19,6 @@ __all__ = ( "CockroachAsyncpgRetryConfig", - "build_connection_config", "build_native_export", "build_native_import", "calculate_backoff_seconds", @@ -31,12 +29,11 @@ "validate_follower_read_staleness", ) +# Retry configuration defaults (module-level for mypyc compatibility) _DEFAULT_MAX_RETRIES: Final[int] = 10 _DEFAULT_BASE_DELAY_MS: Final[float] = 50.0 _DEFAULT_MAX_DELAY_MS: Final[float] = 5000.0 _DEFAULT_ENABLE_LOGGING: Final[bool] = True -_MAX_EXCEPTION_CHAIN_DEPTH: Final[int] = 16 -_RNG: Final = secrets.SystemRandom() @mypyc_attr(allow_interpreted_subclasses=False) @@ -68,26 +65,6 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachAsyncp ) -def build_connection_config(config: "dict[str, Any]") -> "dict[str, Any]": - """Prepare CockroachDB AsyncPG connection config, extracting multi-region server settings.""" - result = asyncpg_build_connection_config(config) - server_settings = dict(result.get("server_settings") or {}) - if "application_name" in result: - server_settings.setdefault("application_name", str(result.pop("application_name"))) - if "gateway_region" in result: - server_settings.setdefault("gateway_region", str(result.pop("gateway_region"))) - if "default_transaction_use_follower_reads" in result: - val = result.pop("default_transaction_use_follower_reads") - server_settings.setdefault( - "default_transaction_use_follower_reads", "on" if val is True or str(val).lower() == "on" else "off" - ) - if "results_buffer_size" in result: - server_settings.setdefault("results_buffer_size", str(result.pop("results_buffer_size"))) - if server_settings: - result["server_settings"] = server_settings - return result - - def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. @@ -108,17 +85,17 @@ def is_retryable_error(error: BaseException) -> bool: Returns: True when the transaction should be retried. """ - depth = 0 + seen: set[int] = set() current: BaseException | None = error - while current is not None and depth < _MAX_EXCEPTION_CHAIN_DEPTH: + while current is not None and id(current) not in seen: + seen.add(id(current)) if isinstance(current, SerializationConflictError): return True if has_sqlstate(current) and str(current.sqlstate) == "40001": return True if not isinstance(current, SQLSpecError): return False - current = cast("BaseException | None", getattr(current, "__cause__", None)) - depth += 1 + current = cast("BaseException | None", cast("Any", current).__cause__) return False @@ -132,7 +109,7 @@ def calculate_backoff_seconds(attempt: int, config: "CockroachAsyncpgRetryConfig capped_ms: float = min(config.base_delay_ms * (2**attempt), config.max_delay_ms) if capped_ms <= 0.0: return 0.0 - return _RNG.uniform(capped_ms / 2.0, capped_ms) / 1000.0 + return random.uniform(capped_ms / 2.0, capped_ms) / 1000.0 # noqa: S311 _STALENESS_LITERAL: Final[re.Pattern[str]] = re.compile(r"'[^'\\;]+'") diff --git a/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py b/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py index 459028332..6839999c3 100644 --- a/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py +++ b/sqlspec/adapters/cockroach_asyncpg/data_dictionary.py @@ -3,7 +3,6 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar -from sqlspec.adapters.asyncpg.data_dictionary import AsyncpgDataDictionary from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -28,9 +27,10 @@ unsupported_system_metadata_capability, ) from sqlspec.data_dictionary.dialects.cockroachdb import resolve_cockroachdb_json_type +from sqlspec.driver import AsyncDataDictionaryBase if TYPE_CHECKING: - from sqlspec.adapters.asyncpg.driver import AsyncpgDriver + from sqlspec.adapters.cockroach_asyncpg.driver import CockroachAsyncpgDriver from sqlspec.core import SQL __all__ = ("CockroachAsyncpgDataDictionary",) @@ -55,19 +55,19 @@ _COCKROACH_SUPPORTED_DOMAINS = frozenset(_COCKROACH_METADATA_DOMAINS) - {"crdb_internal", "system"} -class CockroachAsyncpgDataDictionary(AsyncpgDataDictionary): +class CockroachAsyncpgDataDictionary(AsyncDataDictionaryBase): """CockroachDB async data dictionary (AsyncPG).""" dialect: ClassVar[str] = "cockroachdb" async def get_metadata_capabilities( - self, driver: "AsyncpgDriver", domains: Sequence[str] | None = None + self, driver: "CockroachAsyncpgDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) async def get_system_metadata_capabilities( - self, driver: "AsyncpgDriver", domains: Sequence[str] | None = None + self, driver: "CockroachAsyncpgDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -75,7 +75,7 @@ async def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) async def _select_domain( - self, driver: "AsyncpgDriver", domain: str, query_name: str, **parameters: Any + self, driver: "CockroachAsyncpgDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -83,15 +83,17 @@ async def _select_domain( rows = await driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - async def get_schemas(self, driver: "AsyncpgDriver") -> MetadataResult: + async def get_schemas(self, driver: "CockroachAsyncpgDriver") -> MetadataResult: """Get schema metadata.""" return await self._select_domain(driver, "schemas", "list") - async def get_objects(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_objects(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return await self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_table_details(self, driver: "AsyncpgDriver", table: str, schema: str | None = None) -> MetadataResult: + async def get_table_details( + self, driver: "CockroachAsyncpgDriver", table: str, schema: str | None = None + ) -> MetadataResult: """Get rich table metadata.""" return await self._select_domain( driver, @@ -102,7 +104,7 @@ async def get_table_details(self, driver: "AsyncpgDriver", table: str, schema: s ) async def get_constraints( - self, driver: "AsyncpgDriver", table: str | None = None, schema: str | None = None + self, driver: "CockroachAsyncpgDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -110,16 +112,16 @@ async def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - async def get_views(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_views(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return await self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_routines(self, driver: "AsyncpgDriver", schema: str | None = None) -> MetadataResult: + async def get_routines(self, driver: "CockroachAsyncpgDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) async def get_privileges( - self, driver: "AsyncpgDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachAsyncpgDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -128,7 +130,7 @@ async def get_privileges( ) async def get_dependencies( - self, driver: "AsyncpgDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachAsyncpgDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -138,7 +140,7 @@ async def get_dependencies( async def get_ddl( self, - driver: "AsyncpgDriver", + driver: "CockroachAsyncpgDriver", object_name: str, schema: str | None = None, *, @@ -164,7 +166,7 @@ async def get_ddl( ) async def get_system_metadata( - self, driver: "AsyncpgDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "CockroachAsyncpgDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -172,8 +174,8 @@ async def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - async def get_version(self, driver: "AsyncpgDriver") -> "VersionInfo | None": - driver_id = id(driver.connection) if hasattr(driver, "connection") else id(driver) + async def get_version(self, driver: "CockroachAsyncpgDriver") -> "VersionInfo | None": + driver_id = id(driver) if driver_id in self._version_fetch_attempted: return self._version_cache.get(driver_id) @@ -193,17 +195,17 @@ async def get_version(self, driver: "AsyncpgDriver") -> "VersionInfo | None": self.cache_version(driver_id, version_info) return version_info - async def get_feature_flag(self, driver: "AsyncpgDriver", feature: str) -> bool: + async def get_feature_flag(self, driver: "CockroachAsyncpgDriver", feature: str) -> bool: version_info = await self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - async def get_optimal_type(self, driver: "AsyncpgDriver", type_category: str) -> str: + async def get_optimal_type(self, driver: "CockroachAsyncpgDriver", type_category: str) -> str: config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(await self.get_version(driver)) return config.get_optimal_type(type_category) - async def get_tables(self, driver: "AsyncpgDriver", schema: "str | None" = None) -> "list[TableMetadata]": + async def get_tables(self, driver: "CockroachAsyncpgDriver", schema: "str | None" = None) -> "list[TableMetadata]": schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") return await driver.select( @@ -214,7 +216,7 @@ async def get_tables(self, driver: "AsyncpgDriver", schema: "str | None" = None) ) async def get_columns( - self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": schema_name = self.resolve_schema(schema) if table is None: @@ -237,7 +239,7 @@ async def get_columns( ) async def get_indexes( - self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": schema_name = self.resolve_schema(schema) if table is None: @@ -259,7 +261,7 @@ async def get_indexes( ) async def get_foreign_keys( - self, driver: "AsyncpgDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachAsyncpgDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": schema_name = self.resolve_schema(schema) if table is None: diff --git a/sqlspec/adapters/cockroach_asyncpg/driver.py b/sqlspec/adapters/cockroach_asyncpg/driver.py index 82379cc95..578b8f72e 100644 --- a/sqlspec/adapters/cockroach_asyncpg/driver.py +++ b/sqlspec/adapters/cockroach_asyncpg/driver.py @@ -1,11 +1,10 @@ """CockroachDB AsyncPG driver implementation.""" import asyncio -import contextlib from typing import TYPE_CHECKING, Any, TypeVar, cast from sqlspec.adapters.asyncpg.core import create_mapped_exception, driver_profile -from sqlspec.adapters.asyncpg.driver import AsyncpgDriver, AsyncpgExceptionHandler +from sqlspec.adapters.asyncpg.driver import AsyncpgDriver from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgPostgresError, CockroachAsyncpgSessionContext from sqlspec.adapters.cockroach_asyncpg.core import ( CockroachAsyncpgRetryConfig, @@ -20,6 +19,7 @@ ) from sqlspec.adapters.cockroach_asyncpg.data_dictionary import CockroachAsyncpgDataDictionary from sqlspec.core import SQL, register_driver_profile +from sqlspec.driver import BaseAsyncExceptionHandler from sqlspec.utils.logging import get_logger from sqlspec.utils.type_guards import has_sqlstate @@ -37,7 +37,7 @@ _T = TypeVar("_T") -class CockroachAsyncpgExceptionHandler(AsyncpgExceptionHandler): +class CockroachAsyncpgExceptionHandler(BaseAsyncExceptionHandler): """Async context manager for CockroachDB AsyncPG exceptions.""" __slots__ = () @@ -66,6 +66,7 @@ def __init__( self._retry_config = CockroachAsyncpgRetryConfig.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 select_to_storage( @@ -215,34 +216,61 @@ async def run_transaction_with_retry(self, operation: "Callable[[], Awaitable[_T attempt += 1 async def dispatch_execute(self, cursor: Any, statement: SQL) -> "ExecutionResult": - opened_txn = False - if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): - await self.begin() - opened_txn = True - try: - return await super().dispatch_execute(cursor, statement) - finally: - if opened_txn: - with contextlib.suppress(Exception): - await self.commit() - - def handle_database_exceptions(self) -> "CockroachAsyncpgExceptionHandler": + return await self._dispatch_execute_impl(cursor, statement) + + async def dispatch_execute_many(self, cursor: Any, statement: SQL) -> "ExecutionResult": + return await self._dispatch_execute_many_impl(cursor, statement) + + async def dispatch_execute_script(self, cursor: Any, statement: SQL) -> "ExecutionResult": + return await self._dispatch_execute_script_impl(cursor, statement) + + def handle_database_exceptions(self) -> "CockroachAsyncpgExceptionHandler": # type: ignore[override] return CockroachAsyncpgExceptionHandler() @property - def data_dictionary(self) -> "CockroachAsyncpgDataDictionary": + def data_dictionary(self) -> "CockroachAsyncpgDataDictionary": # type: ignore[override] if self._data_dictionary is None: - self._data_dictionary = CockroachAsyncpgDataDictionary() + # Intentionally assign CockroachDB-specific data dictionary to parent slot + object.__setattr__(self, "_data_dictionary", CockroachAsyncpgDataDictionary()) return cast("CockroachAsyncpgDataDictionary", self._data_dictionary) - def _follower_reads_enabled(self) -> bool: - return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) - async def _apply_follower_reads(self) -> None: - if not self._follower_reads_enabled() or not self._follower_staleness: + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) await self.connection.execute(f"SET TRANSACTION AS OF SYSTEM TIME {staleness}") + async def _begin_follower_read_transaction(self) -> None: + """Open the transaction a follower read needs so the staleness clause can lead it. + + A statement run outside a transaction gets its own implicit one, which + the clause could not precede, so a read opens a transaction here when the + caller has not already done so. + """ + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: + return + if self._connection_in_transaction(): + return + await self.begin() + + async def _dispatch_execute_impl(self, cursor: "CockroachAsyncpgConnection", statement: SQL) -> "ExecutionResult": + if statement.returns_rows(): + await self._begin_follower_read_transaction() + return await super().dispatch_execute(cursor, statement) + + async def _dispatch_execute_many_impl( + self, cursor: "CockroachAsyncpgConnection", statement: SQL + ) -> "ExecutionResult": + return await AsyncpgDriver.dispatch_execute_many(self, cursor, statement) + + async def _dispatch_execute_script_impl( + self, cursor: "CockroachAsyncpgConnection", statement: SQL + ) -> "ExecutionResult": + return await AsyncpgDriver.dispatch_execute_script(self, cursor, statement) + register_driver_profile("cockroach_asyncpg", driver_profile) diff --git a/sqlspec/adapters/cockroach_psycopg/_typing.py b/sqlspec/adapters/cockroach_psycopg/_typing.py index b91c94f7a..eead89a5e 100644 --- a/sqlspec/adapters/cockroach_psycopg/_typing.py +++ b/sqlspec/adapters/cockroach_psycopg/_typing.py @@ -9,9 +9,9 @@ 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.crdb import AsyncCrdbConnection, CrdbConnection from psycopg.rows import DictRow as PsycopgDictRow from psycopg.rows import dict_row as cockroach_psycopg_dict_row from psycopg.types.json import Jsonb as CockroachPsycopgJsonb @@ -23,6 +23,8 @@ from types import TracebackType from typing import TypeAlias + from psycopg.crdb import AsyncCrdbConnection, CrdbConnection + from sqlspec.adapters.cockroach_psycopg.driver import CockroachPsycopgAsyncDriver, CockroachPsycopgSyncDriver from sqlspec.core import StatementConfig @@ -32,8 +34,8 @@ CockroachAsyncCursor: TypeAlias = AsyncCursor[PsycopgDictRow] if not TYPE_CHECKING: - CockroachSyncConnection = CrdbConnection - CockroachAsyncConnection = AsyncCrdbConnection + CockroachSyncConnection = psycopg_crdb.CrdbConnection + CockroachAsyncConnection = 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 58e58f480..8d81ec748 100644 --- a/sqlspec/adapters/cockroach_psycopg/adk/store.py +++ b/sqlspec/adapters/cockroach_psycopg/adk/store.py @@ -8,7 +8,6 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_dict_row as dict_row from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_errors as errors from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_sql as pg_sql -from sqlspec.adapters.cockroach_psycopg.core import as_query from sqlspec.config import ADKConfig from sqlspec.extensions.adk import ( BaseAsyncADKStore, @@ -242,7 +241,7 @@ async def create_session( params = (session_id, app_name, user_id, state_json) async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) row = await cur.fetchone() await conn.commit() @@ -278,7 +277,7 @@ async def get_session( try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (app_name, user_id, session_id)) + await cur.execute(sql.encode(), (app_name, user_id, session_id)) row = await cur.fetchone() if row is None: @@ -303,7 +302,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(as_query(sql), (Jsonb(state), app_name, user_id, session_id)) + await cur.execute(sql.encode(), (Jsonb(state), app_name, user_id, session_id)) await conn.commit() async def list_sessions( @@ -326,7 +325,7 @@ async def list_sessions( try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) rows = await cur.fetchall() except errors.UndefinedTable: return [] @@ -347,7 +346,7 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> sql = f"DELETE FROM {self._session_table} WHERE app_name = %s AND user_id = %s AND id = %s" async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (app_name, user_id, session_id)) + await cur.execute(sql.encode(), (app_name, user_id, session_id)) await conn.commit() async def append_event(self, event_record: StoredEvent) -> None: @@ -361,7 +360,7 @@ async def append_event(self, event_record: StoredEvent) -> None: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: await cur.execute( - as_query(sql), + sql.encode(), ( event_record["id"], event_record["app_name"], @@ -410,7 +409,7 @@ async def append_event_and_update_state( async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: await cur.execute( - as_query(insert_sql), + insert_sql.encode(), ( event_record["id"], event_record["app_name"], @@ -421,14 +420,14 @@ async def append_event_and_update_state( jsonb_value, ), ) - await cur.execute(as_query(update_sql), (Jsonb(state), app_name, user_id, session_id)) + await cur.execute(update_sql.encode(), (Jsonb(state), app_name, user_id, session_id)) row = await cur.fetchone() if row is None: _raise_missing_session(session_id) if app_state is not None: - await cur.execute(as_query(app_upsert_sql), (app_name, Jsonb(app_state))) + await cur.execute(app_upsert_sql.encode(), (app_name, Jsonb(app_state))) if user_state is not None: - await cur.execute(as_query(user_upsert_sql), (app_name, user_id, Jsonb(user_state))) + await cur.execute(user_upsert_sql.encode(), (app_name, user_id, Jsonb(user_state))) except Exception: await conn.rollback() raise @@ -476,7 +475,7 @@ async def get_events( try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), tuple(params)) + await cur.execute(sql.encode(), tuple(params)) rows = await cur.fetchall() except errors.UndefinedTable: return [] @@ -504,7 +503,7 @@ async def delete_expired_events(self, before: "datetime", app_name: "str | None" try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -520,7 +519,7 @@ async def delete_idle_sessions(self, updated_before: "datetime", app_name: "str try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -536,7 +535,7 @@ async def delete_idle_user_states(self, updated_before: "datetime", app_name: "s try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -547,7 +546,7 @@ async def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (app_name,)) + await cur.execute(sql.encode(), (app_name,)) row = await cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -558,7 +557,7 @@ async def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (app_name, user_id)) + await cur.execute(sql.encode(), (app_name, user_id)) row = await cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -571,7 +570,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(as_query(sql), (app_name, Jsonb(state))) + await cur.execute(sql.encode(), (app_name, Jsonb(state))) await conn.commit() async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: @@ -581,7 +580,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(as_query(sql), (app_name, user_id, Jsonb(state))) + await cur.execute(sql.encode(), (app_name, user_id, Jsonb(state))) await conn.commit() async def get_metadata(self, key: str) -> "str | None": @@ -589,7 +588,7 @@ async def get_metadata(self, key: str) -> "str | None": try: async with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (key,)) + await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() return row["value"] if row is not None else None except errors.UndefinedTable: @@ -602,7 +601,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(as_query(sql), (key, value)) + await cur.execute(sql.encode(), (key, value)) await conn.commit() async def _sessions_table_ddl(self) -> str: @@ -728,7 +727,7 @@ def create_session( params = (session_id, app_name, user_id, state_json) with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) row = cur.fetchone() conn.commit() @@ -765,7 +764,7 @@ def get_session( try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name, user_id, session_id)) + cur.execute(sql.encode(), (app_name, user_id, session_id)) row = cur.fetchone() if row is None: @@ -791,7 +790,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(as_query(sql), (Jsonb(state), app_name, user_id, session_id)) + cur.execute(sql.encode(), (Jsonb(state), app_name, user_id, session_id)) conn.commit() def list_sessions( @@ -815,7 +814,7 @@ def list_sessions( try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) rows = cur.fetchall() return [ @@ -837,7 +836,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 = %s AND user_id = %s AND id = %s" with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name, user_id, session_id)) + cur.execute(sql.encode(), (app_name, user_id, session_id)) conn.commit() def append_event(self, event_record: StoredEvent) -> None: @@ -882,7 +881,7 @@ def append_event_and_update_state( with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: try: cur.execute( - as_query(insert_sql), + insert_sql.encode(), ( event_record["id"], event_record["app_name"], @@ -893,14 +892,14 @@ def append_event_and_update_state( jsonb_value, ), ) - cur.execute(as_query(update_sql), (Jsonb(state), app_name, user_id, session_id)) + cur.execute(update_sql.encode(), (Jsonb(state), app_name, user_id, session_id)) row = cur.fetchone() if row is None: _raise_missing_session(session_id) if app_state is not None: - cur.execute(as_query(app_upsert_sql), (app_name, Jsonb(app_state))) + cur.execute(app_upsert_sql.encode(), (app_name, Jsonb(app_state))) if user_state is not None: - cur.execute(as_query(user_upsert_sql), (app_name, user_id, Jsonb(user_state))) + cur.execute(user_upsert_sql.encode(), (app_name, user_id, Jsonb(user_state))) except Exception: conn.rollback() raise @@ -949,7 +948,7 @@ def get_events( try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), tuple(params)) + cur.execute(sql.encode(), tuple(params)) rows = cur.fetchall() return [ @@ -978,7 +977,7 @@ def delete_expired_events(self, before: "datetime", app_name: "str | None" = Non try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -995,7 +994,7 @@ def delete_idle_sessions(self, updated_before: "datetime", app_name: "str | None try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -1012,7 +1011,7 @@ def delete_idle_user_states(self, updated_before: "datetime", app_name: "str | N try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 except errors.UndefinedTable: @@ -1024,7 +1023,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name,)) + cur.execute(sql.encode(), (app_name,)) row = cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -1036,7 +1035,7 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name, user_id)) + cur.execute(sql.encode(), (app_name, user_id)) row = cur.fetchone() return row["state"] if row is not None else None except errors.UndefinedTable: @@ -1050,7 +1049,7 @@ def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name, Jsonb(state))) + cur.execute(sql.encode(), (app_name, Jsonb(state))) conn.commit() def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: @@ -1061,7 +1060,7 @@ def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]" """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (app_name, user_id, Jsonb(state))) + cur.execute(sql.encode(), (app_name, user_id, Jsonb(state))) conn.commit() def get_metadata(self, key: str) -> "str | None": @@ -1070,7 +1069,7 @@ def get_metadata(self, key: str) -> "str | None": try: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (key,)) + cur.execute(sql.encode(), (key,)) row = cur.fetchone() return row["value"] if row is not None else None except errors.UndefinedTable: @@ -1084,7 +1083,7 @@ def set_metadata(self, key: str, value: str) -> None: """ with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (key, value)) + cur.execute(sql.encode(), (key, value)) conn.commit() def _sessions_table_ddl(self) -> str: @@ -1177,7 +1176,7 @@ def _insert_event(self, event_record: StoredEvent) -> None: with self._config.provide_connection() as conn, conn.cursor(row_factory=dict_row) as cur: cur.execute( - as_query(sql), + sql.encode(), ( event_record["id"], event_record["app_name"], @@ -1296,7 +1295,7 @@ async def search_entries( try: async with self._config.provide_connection() as conn, conn.cursor() as cur: - await cur.execute(as_query(sql), params) + await cur.execute(sql.encode(), params) rows = await cur.fetchall() columns = [col[0] for col in cur.description or []] except errors.UndefinedTable: @@ -1316,7 +1315,7 @@ async def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = %s" async with self._config.provide_connection() as conn, conn.cursor() as cur: - await cur.execute(as_query(sql), (session_id,)) + await cur.execute(sql.encode(), (session_id,)) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1338,7 +1337,7 @@ 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, conn.cursor() as cur: - await cur.execute(as_query(sql), tuple(params) if params else None) + await cur.execute(sql.encode(), tuple(params) if params else None) await conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1480,7 +1479,7 @@ def search_entries( try: with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(as_query(sql), params) + cur.execute(sql.encode(), params) rows = cur.fetchall() columns = [col[0] for col in cur.description or []] except errors.UndefinedTable: @@ -1501,7 +1500,7 @@ def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = %s" with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(as_query(sql), (session_id,)) + cur.execute(sql.encode(), (session_id,)) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 @@ -1522,7 +1521,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, conn.cursor() as cur: - cur.execute(as_query(sql), tuple(params) if params else None) + cur.execute(sql.encode(), tuple(params) if params else None) conn.commit() return cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 diff --git a/sqlspec/adapters/cockroach_psycopg/config.py b/sqlspec/adapters/cockroach_psycopg/config.py index 6e04ebcef..58b9c1ac5 100644 --- a/sqlspec/adapters/cockroach_psycopg/config.py +++ b/sqlspec/adapters/cockroach_psycopg/config.py @@ -17,7 +17,6 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_crdb as psycopg_crdb from sqlspec.adapters.cockroach_psycopg.core import ( apply_driver_features, - build_connection_config, build_statement_config, validate_follower_read_staleness, ) @@ -38,9 +37,10 @@ ) from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints +from sqlspec.utils.config_tools import normalize_connection_config if TYPE_CHECKING: - from collections.abc import Awaitable, Callable + from collections.abc import Awaitable, Callable, Mapping from types import TracebackType from sqlspec.core import StatementConfig @@ -70,11 +70,6 @@ class CockroachPsycopgConnectionConfig(TypedDict): connect_timeout: NotRequired[int] options: NotRequired[str] application_name: NotRequired[str] - gateway_region: NotRequired[str] - default_transaction_use_follower_reads: NotRequired[bool] - results_buffer_size: NotRequired[int] - statement_timeout: NotRequired[int] - idle_in_transaction_session_timeout: NotRequired[int] sslmode: NotRequired[str] sslcert: NotRequired[str] sslkey: NotRequired[str] @@ -147,6 +142,33 @@ class CockroachPsycopgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["poll_queue"]] +def build_connection_config( + connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", +) -> dict[str, Any]: + """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility. + + Maps connection string aliases (dsn, url, connection_string) to conninfo, database aliases + (database, db) to dbname, and user aliases (username) to user, while discarding redundant keys + that libpq rejects. + """ + config = normalize_connection_config(connection_config) + conninfo = ( + config.pop("conninfo", None) + or config.pop("dsn", None) + or config.pop("url", None) + or config.pop("connection_string", None) + ) + if conninfo is not None: + config["conninfo"] = conninfo + dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) + if dbname is not None: + config["dbname"] = dbname + user = config.pop("user", None) or config.pop("username", None) + if user is not None: + config["user"] = user + return config + + class CockroachPsycopgSyncConnectionContext(SyncPoolConnectionContext): """Context manager for CockroachDB psycopg connections.""" @@ -255,7 +277,7 @@ def _create_pool(self) -> "ConnectionPool": "name": all_config.pop("name", None), "timeout": all_config.pop("timeout", 30.0), "max_waiting": all_config.pop("max_waiting", 0), - "max_lifetime": all_config.pop("max_lifetime", 1800.0), + "max_lifetime": all_config.pop("max_lifetime", 3600.0), "max_idle": all_config.pop("max_idle", 600.0), "reconnect_timeout": all_config.pop("reconnect_timeout", 300.0), "reconnect_failed": all_config.pop("reconnect_failed", None), @@ -285,6 +307,7 @@ def _configure_connection(self, conn: "CockroachSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting + # Call user-provided callback after internal setup if self._user_connection_hook is not None: self._user_connection_hook(conn) @@ -481,7 +504,7 @@ async def _create_pool(self) -> "AsyncConnectionPool": "name": all_config.pop("name", None), "timeout": all_config.pop("timeout", 30.0), "max_waiting": all_config.pop("max_waiting", 0), - "max_lifetime": all_config.pop("max_lifetime", 1800.0), + "max_lifetime": all_config.pop("max_lifetime", 3600.0), "max_idle": all_config.pop("max_idle", 600.0), "reconnect_timeout": all_config.pop("reconnect_timeout", 300.0), "reconnect_failed": all_config.pop("reconnect_failed", None), @@ -518,6 +541,7 @@ async def _configure_async_connection(self, conn: "CockroachAsyncConnection") -> if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) + # 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/cockroach_psycopg/core.py b/sqlspec/adapters/cockroach_psycopg/core.py index e873ac1f7..19dbc8257 100644 --- a/sqlspec/adapters/cockroach_psycopg/core.py +++ b/sqlspec/adapters/cockroach_psycopg/core.py @@ -1,31 +1,26 @@ """CockroachDB psycopg adapter compiled helpers.""" +import random import re -import secrets from typing import TYPE_CHECKING, Any, Final, cast from mypy_extensions import mypyc_attr from sqlglot import tokenize from sqlglot.tokenizer_core import TokenType -from typing_extensions import LiteralString from sqlspec.adapters.psycopg.core import apply_driver_features, build_statement_config, driver_profile from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError -from sqlspec.utils.config_tools import normalize_connection_config from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate if TYPE_CHECKING: from collections.abc import Mapping - from sqlspec.adapters.cockroach_psycopg.config import CockroachPsycopgPoolConfig from sqlspec.storage import StorageTelemetry __all__ = ( "CockroachPsycopgRetryConfig", "apply_driver_features", - "as_query", - "build_connection_config", "build_native_export", "build_native_import", "build_statement_config", @@ -38,14 +33,14 @@ "validate_follower_read_staleness", ) +# Retry configuration defaults (module-level for mypyc compatibility) _DEFAULT_MAX_RETRIES: Final[int] = 10 _DEFAULT_BASE_DELAY_MS: Final[float] = 50.0 _DEFAULT_MAX_DELAY_MS: Final[float] = 5000.0 _DEFAULT_ENABLE_LOGGING: Final[bool] = True -_MAX_EXCEPTION_CHAIN_DEPTH: Final[int] = 16 -_RNG: Final = secrets.SystemRandom() +# Keep this in sync with cockroach_asyncpg.core.CockroachAsyncpgRetryConfig. @mypyc_attr(allow_interpreted_subclasses=False) class CockroachPsycopgRetryConfig: """CockroachDB psycopg transaction retry configuration.""" @@ -75,49 +70,6 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachPsycop ) -def build_connection_config( - connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", -) -> dict[str, Any]: - """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility.""" - config = normalize_connection_config(connection_config) - conninfo = ( - config.pop("conninfo", None) - or config.pop("dsn", None) - or config.pop("url", None) - or config.pop("connection_string", None) - ) - if conninfo is not None: - config["conninfo"] = conninfo - dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) - if dbname is not None: - config["dbname"] = dbname - user = config.pop("user", None) or config.pop("username", None) - if user is not None: - config["user"] = user - - session_options: list[str] = [] - if "gateway_region" in config: - session_options.append(f"-c gateway_region={config.pop('gateway_region')}") - if "default_transaction_use_follower_reads" in config: - val = config.pop("default_transaction_use_follower_reads") - val_str = "on" if val is True or str(val).lower() == "on" else "off" - session_options.append(f"-c default_transaction_use_follower_reads={val_str}") - if "results_buffer_size" in config: - session_options.append(f"-c results_buffer_size={config.pop('results_buffer_size')}") - if "statement_timeout" in config: - session_options.append(f"-c statement_timeout={config.pop('statement_timeout')}") - if "idle_in_transaction_session_timeout" in config: - session_options.append( - f"-c idle_in_transaction_session_timeout={config.pop('idle_in_transaction_session_timeout')}" - ) - if session_options: - existing = config.get("options") - opt_str = " ".join(session_options) - config["options"] = f"{existing} {opt_str}" if existing else opt_str - - return config - - def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. @@ -138,17 +90,17 @@ def is_retryable_error(error: BaseException) -> bool: Returns: True when the transaction should be retried. """ - depth = 0 + seen: set[int] = set() current: BaseException | None = error - while current is not None and depth < _MAX_EXCEPTION_CHAIN_DEPTH: + while current is not None and id(current) not in seen: + seen.add(id(current)) if isinstance(current, SerializationConflictError): return True if has_sqlstate(current) and str(current.sqlstate) == "40001": return True if not isinstance(current, SQLSpecError): return False - current = cast("BaseException | None", getattr(current, "__cause__", None)) - depth += 1 + current = cast("BaseException | None", cast("Any", current).__cause__) return False @@ -162,19 +114,7 @@ def calculate_backoff_seconds(attempt: int, config: "CockroachPsycopgRetryConfig capped_ms: float = min(config.base_delay_ms * (2**attempt), config.max_delay_ms) if capped_ms <= 0.0: return 0.0 - return _RNG.uniform(capped_ms / 2.0, capped_ms) / 1000.0 - - -def as_query(sql: object) -> LiteralString: - """Prepare a SQL string for psycopg query execution without byte encoding. - - Args: - sql: The raw SQL query string or object. - - Returns: - The SQL query string typed as a LiteralString for driver query dispatch. - """ - return cast("LiteralString", sql) + return random.uniform(capped_ms / 2.0, capped_ms) / 1000.0 # noqa: S311 _STALENESS_LITERAL: Final[re.Pattern[str]] = re.compile(r"'[^'\\;]+'") diff --git a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py index 6f525d864..b9a8a78ad 100644 --- a/sqlspec/adapters/cockroach_psycopg/data_dictionary.py +++ b/sqlspec/adapters/cockroach_psycopg/data_dictionary.py @@ -3,7 +3,8 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar -from sqlspec.adapters.psycopg.data_dictionary import PsycopgAsyncDataDictionary, PsycopgSyncDataDictionary +from mypy_extensions import mypyc_attr + from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -28,9 +29,10 @@ unsupported_system_metadata_capability, ) from sqlspec.data_dictionary.dialects.cockroachdb import resolve_cockroachdb_json_type +from sqlspec.driver import AsyncDataDictionaryBase, SyncDataDictionaryBase if TYPE_CHECKING: - from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver + from sqlspec.adapters.cockroach_psycopg.driver import CockroachPsycopgAsyncDriver, CockroachPsycopgSyncDriver from sqlspec.core import SQL __all__ = ("CockroachPsycopgAsyncDataDictionary", "CockroachPsycopgSyncDataDictionary") @@ -55,7 +57,8 @@ _COCKROACH_SUPPORTED_DOMAINS = frozenset(_COCKROACH_METADATA_DOMAINS) - {"crdb_internal", "system"} -class CockroachPsycopgSyncDataDictionary(PsycopgSyncDataDictionary): +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) +class CockroachPsycopgSyncDataDictionary(SyncDataDictionaryBase): """CockroachDB sync data dictionary.""" dialect: ClassVar[str] = "cockroachdb" @@ -64,13 +67,13 @@ def __init__(self) -> None: super().__init__() def get_metadata_capabilities( - self, driver: "PsycopgSyncDriver", domains: Sequence[str] | None = None + self, driver: "CockroachPsycopgSyncDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) def get_system_metadata_capabilities( - self, driver: "PsycopgSyncDriver", domains: Sequence[str] | None = None + self, driver: "CockroachPsycopgSyncDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -78,7 +81,7 @@ def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) def _select_domain( - self, driver: "PsycopgSyncDriver", domain: str, query_name: str, **parameters: Any + self, driver: "CockroachPsycopgSyncDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -86,15 +89,17 @@ def _select_domain( rows = driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - def get_schemas(self, driver: "PsycopgSyncDriver") -> MetadataResult: + def get_schemas(self, driver: "CockroachPsycopgSyncDriver") -> MetadataResult: """Get schema metadata.""" return self._select_domain(driver, "schemas", "list") - def get_objects(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_objects(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) - def get_table_details(self, driver: "PsycopgSyncDriver", table: str, schema: str | None = None) -> MetadataResult: + def get_table_details( + self, driver: "CockroachPsycopgSyncDriver", table: str, schema: str | None = None + ) -> MetadataResult: """Get rich table metadata.""" return self._select_domain( driver, @@ -105,7 +110,7 @@ def get_table_details(self, driver: "PsycopgSyncDriver", table: str, schema: str ) def get_constraints( - self, driver: "PsycopgSyncDriver", table: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgSyncDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -113,16 +118,16 @@ def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - def get_views(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_views(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - def get_routines(self, driver: "PsycopgSyncDriver", schema: str | None = None) -> MetadataResult: + def get_routines(self, driver: "CockroachPsycopgSyncDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) def get_privileges( - self, driver: "PsycopgSyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgSyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -131,7 +136,7 @@ def get_privileges( ) def get_dependencies( - self, driver: "PsycopgSyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgSyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -141,7 +146,7 @@ def get_dependencies( def get_ddl( self, - driver: "PsycopgSyncDriver", + driver: "CockroachPsycopgSyncDriver", object_name: str, schema: str | None = None, *, @@ -167,7 +172,7 @@ def get_ddl( ) def get_system_metadata( - self, driver: "PsycopgSyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "CockroachPsycopgSyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -175,41 +180,41 @@ def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - def get_version(self, driver: "PsycopgSyncDriver") -> "VersionInfo | None": + def get_version(self, driver: "CockroachPsycopgSyncDriver") -> "VersionInfo | None": """Get CockroachDB version information.""" - cache_key = id(driver.connection) - if cache_key in self._version_fetch_attempted: - return self._version_cache.get(cache_key) + driver_id = id(driver) + if driver_id in self._version_fetch_attempted: + return self._version_cache.get(driver_id) version_value = driver.select_value_or_none(self.get_query("version", "current")) if not version_value: self._log_version_unavailable(type(self).dialect, "missing") - self.cache_version(cache_key, None) + self.cache_version(driver_id, None) return None version_info = self.parse_version_with_pattern(self.get_dialect_config().version_pattern, str(version_value)) if version_info is None: self._log_version_unavailable(type(self).dialect, "parse_failed") - self.cache_version(cache_key, None) + self.cache_version(driver_id, None) return None self._log_version_detected(type(self).dialect, version_info) - self.cache_version(cache_key, version_info) + self.cache_version(driver_id, version_info) return version_info - def get_feature_flag(self, driver: "PsycopgSyncDriver", feature: str) -> bool: + def get_feature_flag(self, driver: "CockroachPsycopgSyncDriver", feature: str) -> bool: """Check if CockroachDB supports a specific feature.""" version_info = self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - def get_optimal_type(self, driver: "PsycopgSyncDriver", type_category: str) -> str: + def get_optimal_type(self, driver: "CockroachPsycopgSyncDriver", type_category: str) -> str: """Get optimal CockroachDB type for a category.""" config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(self.get_version(driver)) return config.get_optimal_type(type_category) - def get_tables(self, driver: "PsycopgSyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": + def get_tables(self, driver: "CockroachPsycopgSyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": """Get tables sorted by dependency order.""" schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -221,7 +226,7 @@ def get_tables(self, driver: "PsycopgSyncDriver", schema: "str | None" = None) - ) def get_columns( - self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": """Get column information for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -245,7 +250,7 @@ def get_columns( ) def get_indexes( - self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": """Get index metadata for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -268,7 +273,7 @@ def get_indexes( ) def get_foreign_keys( - self, driver: "PsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgSyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": """Get foreign key metadata.""" schema_name = self.resolve_schema(schema) @@ -288,7 +293,8 @@ def get_foreign_keys( ) -class CockroachPsycopgAsyncDataDictionary(PsycopgAsyncDataDictionary): +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) +class CockroachPsycopgAsyncDataDictionary(AsyncDataDictionaryBase): """CockroachDB async data dictionary.""" dialect: ClassVar[str] = "cockroachdb" @@ -297,13 +303,13 @@ def __init__(self) -> None: super().__init__() async def get_metadata_capabilities( - self, driver: "PsycopgAsyncDriver", domains: Sequence[str] | None = None + self, driver: "CockroachPsycopgAsyncDriver", domains: Sequence[str] | None = None ) -> MetadataCapabilityProfile: """Get CockroachDB data-dictionary capability profile.""" return _cockroach_metadata_profile(type(self).__name__, domains) async def get_system_metadata_capabilities( - self, driver: "PsycopgAsyncDriver", domains: Sequence[str] | None = None + self, driver: "CockroachPsycopgAsyncDriver", domains: Sequence[str] | None = None ) -> tuple[SystemMetadataCapability, ...]: """Get CockroachDB opt-in system metadata capability disclosures.""" _ = driver @@ -311,7 +317,7 @@ async def get_system_metadata_capabilities( return tuple(unsupported_system_metadata_capability(domain) for domain in requested_domains) async def _select_domain( - self, driver: "PsycopgAsyncDriver", domain: str, query_name: str, **parameters: Any + self, driver: "CockroachPsycopgAsyncDriver", domain: str, query_name: str, **parameters: Any ) -> MetadataResult: query = get_data_dictionary_loader().get_domain_query(type(self).dialect, domain, query_name) if not query.is_supported or query.sql is None: @@ -319,16 +325,16 @@ async def _select_domain( rows = await driver.select(query.sql, **parameters) return _metadata_result(domain, _cockroach_metadata_capability(domain), rows) - async def get_schemas(self, driver: "PsycopgAsyncDriver") -> MetadataResult: + async def get_schemas(self, driver: "CockroachPsycopgAsyncDriver") -> MetadataResult: """Get schema metadata.""" return await self._select_domain(driver, "schemas", "list") - async def get_objects(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_objects(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get database object metadata.""" return await self._select_domain(driver, "objects", "by_schema", schema_name=self.resolve_schema(schema)) async def get_table_details( - self, driver: "PsycopgAsyncDriver", table: str, schema: str | None = None + self, driver: "CockroachPsycopgAsyncDriver", table: str, schema: str | None = None ) -> MetadataResult: """Get rich table metadata.""" return await self._select_domain( @@ -340,7 +346,7 @@ async def get_table_details( ) async def get_constraints( - self, driver: "PsycopgAsyncDriver", table: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgAsyncDriver", table: str | None = None, schema: str | None = None ) -> MetadataResult: """Get constraint metadata.""" table_name = self.resolve_identifier(table) if table is not None else None @@ -348,16 +354,16 @@ async def get_constraints( driver, "constraints", "by_schema", schema_name=self.resolve_schema(schema), table_name=table_name ) - async def get_views(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_views(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get view metadata.""" return await self._select_domain(driver, "views", "by_schema", schema_name=self.resolve_schema(schema)) - async def get_routines(self, driver: "PsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: + async def get_routines(self, driver: "CockroachPsycopgAsyncDriver", schema: str | None = None) -> MetadataResult: """Get routine metadata.""" return _metadata_result("routines", MetadataCapability.unsupported("routines")) async def get_privileges( - self, driver: "PsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get privilege metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -366,7 +372,7 @@ async def get_privileges( ) async def get_dependencies( - self, driver: "PsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None + self, driver: "CockroachPsycopgAsyncDriver", object_name: str | None = None, schema: str | None = None ) -> MetadataResult: """Get stable dependency metadata.""" resolved_object = self.resolve_identifier(object_name) if object_name is not None else None @@ -376,7 +382,7 @@ async def get_dependencies( async def get_ddl( self, - driver: "PsycopgAsyncDriver", + driver: "CockroachPsycopgAsyncDriver", object_name: str, schema: str | None = None, *, @@ -402,7 +408,7 @@ async def get_ddl( ) async def get_system_metadata( - self, driver: "PsycopgAsyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any + self, driver: "CockroachPsycopgAsyncDriver", request: SystemMetadataRequest | str | None = None, **kwargs: Any ) -> SystemMetadataResult: """Get opt-in CockroachDB system metadata.""" _ = driver @@ -410,41 +416,43 @@ async def get_system_metadata( capability = unsupported_system_metadata_capability(metadata_request.domain) return system_metadata_gated_result(metadata_request, capability) - async def get_version(self, driver: "PsycopgAsyncDriver") -> "VersionInfo | None": + async def get_version(self, driver: "CockroachPsycopgAsyncDriver") -> "VersionInfo | None": """Get CockroachDB version information.""" - cache_key = id(driver.connection) - if cache_key in self._version_fetch_attempted: - return self._version_cache.get(cache_key) + driver_id = id(driver) + if driver_id in self._version_fetch_attempted: + return self._version_cache.get(driver_id) version_value = await driver.select_value_or_none(self.get_query("version", "current")) if not version_value: self._log_version_unavailable(type(self).dialect, "missing") - self.cache_version(cache_key, None) + self.cache_version(driver_id, None) return None version_info = self.parse_version_with_pattern(self.get_dialect_config().version_pattern, str(version_value)) if version_info is None: self._log_version_unavailable(type(self).dialect, "parse_failed") - self.cache_version(cache_key, None) + self.cache_version(driver_id, None) return None self._log_version_detected(type(self).dialect, version_info) - self.cache_version(cache_key, version_info) + self.cache_version(driver_id, version_info) return version_info - async def get_feature_flag(self, driver: "PsycopgAsyncDriver", feature: str) -> bool: + async def get_feature_flag(self, driver: "CockroachPsycopgAsyncDriver", feature: str) -> bool: """Check if CockroachDB supports a specific feature.""" version_info = await self.get_version(driver) return self.resolve_feature_flag(feature, version_info) - async def get_optimal_type(self, driver: "PsycopgAsyncDriver", type_category: str) -> str: + async def get_optimal_type(self, driver: "CockroachPsycopgAsyncDriver", type_category: str) -> str: """Get optimal CockroachDB type for a category.""" config = self.get_dialect_config() if type_category == "json": return resolve_cockroachdb_json_type(await self.get_version(driver)) return config.get_optimal_type(type_category) - async def get_tables(self, driver: "PsycopgAsyncDriver", schema: "str | None" = None) -> "list[TableMetadata]": + async def get_tables( + self, driver: "CockroachPsycopgAsyncDriver", schema: "str | None" = None + ) -> "list[TableMetadata]": """Get tables sorted by dependency order.""" schema_name = self.resolve_schema(schema) self._log_schema_introspect(driver, schema_name=schema_name, table_name=None, operation="tables") @@ -456,7 +464,7 @@ async def get_tables(self, driver: "PsycopgAsyncDriver", schema: "str | None" = ) async def get_columns( - self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ColumnMetadata]": """Get column information for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -480,7 +488,7 @@ async def get_columns( ) async def get_indexes( - self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[IndexMetadata]": """Get index metadata for a table or schema.""" schema_name = self.resolve_schema(schema) @@ -503,7 +511,7 @@ async def get_indexes( ) async def get_foreign_keys( - self, driver: "PsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None + self, driver: "CockroachPsycopgAsyncDriver", table: "str | None" = None, schema: "str | None" = None ) -> "list[ForeignKeyMetadata]": """Get foreign key metadata.""" schema_name = self.resolve_schema(schema) diff --git a/sqlspec/adapters/cockroach_psycopg/driver.py b/sqlspec/adapters/cockroach_psycopg/driver.py index cd4be1209..7871e1527 100644 --- a/sqlspec/adapters/cockroach_psycopg/driver.py +++ b/sqlspec/adapters/cockroach_psycopg/driver.py @@ -1,7 +1,6 @@ """CockroachDB psycopg driver implementation.""" import asyncio -import contextlib import time from typing import TYPE_CHECKING, Any, TypeVar, cast @@ -15,7 +14,6 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_module as psycopg from sqlspec.adapters.cockroach_psycopg.core import ( CockroachPsycopgRetryConfig, - as_query, build_native_export, build_native_import, build_statement_config, @@ -32,13 +30,9 @@ CockroachPsycopgSyncDataDictionary, ) from sqlspec.adapters.psycopg.core import create_mapped_exception -from sqlspec.adapters.psycopg.driver import ( - PsycopgAsyncDriver, - PsycopgAsyncExceptionHandler, - PsycopgSyncDriver, - PsycopgSyncExceptionHandler, -) +from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.driver import BaseAsyncExceptionHandler, BaseSyncExceptionHandler from sqlspec.utils.logging import get_logger if TYPE_CHECKING: @@ -61,7 +55,7 @@ _T = TypeVar("_T") -class CockroachPsycopgSyncExceptionHandler(PsycopgSyncExceptionHandler): +class CockroachPsycopgSyncExceptionHandler(BaseSyncExceptionHandler): """Context manager for handling CockroachDB psycopg exceptions.""" __slots__ = () @@ -75,7 +69,7 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False -class CockroachPsycopgAsyncExceptionHandler(PsycopgAsyncExceptionHandler): +class CockroachPsycopgAsyncExceptionHandler(BaseAsyncExceptionHandler): """Async context manager for handling CockroachDB psycopg exceptions.""" __slots__ = () @@ -111,6 +105,7 @@ 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 select_to_storage( @@ -213,7 +208,7 @@ def _execute_native_storage(self, command: str, parameters: "list[Any]") -> "lis rows = [] with self.with_cursor(self.connection) as cursor, handler: cursor.row_factory = dict_row - cursor.execute(as_query(command), parameters) + cursor.execute(command.encode("utf-8"), parameters) rows = cursor.fetchall() if handler.pending_exception is not None: raise handler.pending_exception @@ -266,35 +261,58 @@ def run_transaction_with_retry(self, operation: "Callable[[], _T]") -> _T: attempt += 1 def dispatch_execute(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": - opened_txn = False - if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): - self.begin() - opened_txn = True - try: - return super().dispatch_execute(cursor, statement) - finally: - if opened_txn: - with contextlib.suppress(Exception): - self.commit() - - def handle_database_exceptions(self) -> "CockroachPsycopgSyncExceptionHandler": + return self._dispatch_execute_impl(cursor, statement) + + def dispatch_execute_many(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": + return self._dispatch_execute_many_impl(cursor, statement) + + def dispatch_execute_script(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": + return self._dispatch_execute_script_impl(cursor, statement) + + def handle_database_exceptions(self) -> "CockroachPsycopgSyncExceptionHandler": # type: ignore[override] return CockroachPsycopgSyncExceptionHandler() @property - def data_dictionary(self) -> "CockroachPsycopgSyncDataDictionary": + def data_dictionary(self) -> "CockroachPsycopgSyncDataDictionary": # type: ignore[override] if self._data_dictionary is None: - self._data_dictionary = CockroachPsycopgSyncDataDictionary() + # Intentionally assign CockroachDB-specific data dictionary to parent slot + self._data_dictionary = CockroachPsycopgSyncDataDictionary() # type: ignore[assignment] return cast("CockroachPsycopgSyncDataDictionary", self._data_dictionary) - def _follower_reads_enabled(self) -> bool: - return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) - def _apply_follower_reads(self) -> None: - if not self._follower_reads_enabled() or not self._follower_staleness: + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) self.connection.execute(cast("Any", f"SET TRANSACTION AS OF SYSTEM TIME {staleness}")).close() + def _begin_follower_read_transaction(self) -> None: + """Open the transaction a follower read needs so the staleness clause can lead it. + + psycopg opens a transaction on the first statement, which would leave the + clause with nowhere to go, so a read opens one here when the caller has + not already done so. + """ + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: + return + if self._connection_in_transaction(): + return + self.begin() + + def _dispatch_execute_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": + if statement.returns_rows(): + self._begin_follower_read_transaction() + return super().dispatch_execute(cursor, statement) + + def _dispatch_execute_many_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": + return PsycopgSyncDriver.dispatch_execute_many(self, cursor, statement) + + def _dispatch_execute_script_impl(self, cursor: "CockroachSyncCursor", statement: SQL) -> "ExecutionResult": + return PsycopgSyncDriver.dispatch_execute_script(self, cursor, statement) + class CockroachPsycopgAsyncDriver(PsycopgAsyncDriver): """CockroachDB async driver using psycopg.crdb.""" @@ -318,6 +336,7 @@ 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 select_to_storage( @@ -420,7 +439,7 @@ async def _execute_native_storage(self, command: str, parameters: "list[Any]") - rows = [] async with self.with_cursor(self.connection) as cursor, handler: cursor.row_factory = dict_row - await cursor.execute(as_query(command), parameters) + await cursor.execute(command.encode("utf-8"), parameters) rows = await cursor.fetchall() if handler.pending_exception is not None: raise handler.pending_exception @@ -473,35 +492,58 @@ async def run_transaction_with_retry(self, operation: "Callable[[], Awaitable[_T attempt += 1 async def dispatch_execute(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": - opened_txn = False - if statement.returns_rows() and not self._connection_in_transaction() and self._follower_reads_enabled(): - await self.begin() - opened_txn = True - try: - return await super().dispatch_execute(cursor, statement) - finally: - if opened_txn: - with contextlib.suppress(Exception): - await self.commit() - - def handle_database_exceptions(self) -> "CockroachPsycopgAsyncExceptionHandler": + return await self._dispatch_execute_impl(cursor, statement) + + async def dispatch_execute_many(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": + return await self._dispatch_execute_many_impl(cursor, statement) + + async def dispatch_execute_script(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": + return await self._dispatch_execute_script_impl(cursor, statement) + + def handle_database_exceptions(self) -> "CockroachPsycopgAsyncExceptionHandler": # type: ignore[override] return CockroachPsycopgAsyncExceptionHandler() @property - def data_dictionary(self) -> "CockroachPsycopgAsyncDataDictionary": + def data_dictionary(self) -> "CockroachPsycopgAsyncDataDictionary": # type: ignore[override] if self._data_dictionary is None: - self._data_dictionary = CockroachPsycopgAsyncDataDictionary() + # Intentionally assign CockroachDB-specific data dictionary to parent slot + self._data_dictionary = CockroachPsycopgAsyncDataDictionary() # type: ignore[assignment] return cast("CockroachPsycopgAsyncDataDictionary", self._data_dictionary) - def _follower_reads_enabled(self) -> bool: - return bool(self.driver_features.get("enable_follower_reads", False) and self._follower_staleness) - async def _apply_follower_reads(self) -> None: - if not self._follower_reads_enabled() or not self._follower_staleness: + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: return staleness = validate_follower_read_staleness(self._follower_staleness) cursor = await self.connection.execute(cast("Any", f"SET TRANSACTION AS OF SYSTEM TIME {staleness}")) await cursor.close() + async def _begin_follower_read_transaction(self) -> None: + """Open the transaction a follower read needs so the staleness clause can lead it. + + psycopg opens a transaction on the first statement, which would leave the + clause with nowhere to go, so a read opens one here when the caller has + not already done so. + """ + if not self.driver_features.get("enable_follower_reads", False): + return + if not self._follower_staleness: + return + if self._connection_in_transaction(): + return + await self.begin() + + async def _dispatch_execute_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": + if statement.returns_rows(): + await self._begin_follower_read_transaction() + return await super().dispatch_execute(cursor, statement) + + async def _dispatch_execute_many_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": + return await PsycopgAsyncDriver.dispatch_execute_many(self, cursor, statement) + + async def _dispatch_execute_script_impl(self, cursor: "CockroachAsyncCursor", statement: SQL) -> "ExecutionResult": + return await PsycopgAsyncDriver.dispatch_execute_script(self, cursor, statement) + register_driver_profile("cockroach_psycopg", driver_profile) diff --git a/sqlspec/adapters/cockroach_psycopg/litestar/store.py b/sqlspec/adapters/cockroach_psycopg/litestar/store.py index c7bbad239..8b915ca1e 100644 --- a/sqlspec/adapters/cockroach_psycopg/litestar/store.py +++ b/sqlspec/adapters/cockroach_psycopg/litestar/store.py @@ -6,7 +6,6 @@ from typing_extensions import NotRequired from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_dict_row as dict_row -from sqlspec.adapters.cockroach_psycopg.core import as_query from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -68,7 +67,7 @@ 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: async with conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (key,)) + await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() if row is None: @@ -82,7 +81,7 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by SET expires_at = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ - await conn.execute(as_query(update_sql), (new_expires_at, key)) + await conn.execute(update_sql.encode(), (new_expires_at, key)) await conn.commit() return bytes(row["data"]) @@ -103,7 +102,7 @@ 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: - await conn.execute(as_query(sql), (key, data, expires_at)) + await conn.execute(sql.encode(), (key, data, expires_at)) await conn.commit() async def delete(self, key: str) -> None: @@ -111,7 +110,7 @@ async def delete(self, key: str) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - await conn.execute(as_query(sql), (key,)) + await conn.execute(sql.encode(), (key,)) await conn.commit() async def delete_all(self) -> None: @@ -119,7 +118,7 @@ async def delete_all(self) -> None: conn_context = self._config.provide_connection() async with conn_context as conn: - await conn.execute(as_query(sql)) + await conn.execute(sql.encode()) await conn.commit() self._log_delete_all() @@ -132,7 +131,7 @@ async def exists(self, key: str) -> bool: conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: - await cur.execute(as_query(sql), (key,)) + await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() return row is not None @@ -145,7 +144,7 @@ async def expires_in(self, key: str) -> "int | None": conn_context = self._config.provide_connection() async with conn_context as conn: async with conn.cursor(row_factory=dict_row) as cur: - await cur.execute(as_query(sql), (key,)) + await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() if row is None or row["expires_at"] is None: @@ -165,7 +164,7 @@ async def delete_expired(self) -> int: conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: - await cur.execute(as_query(sql)) + await cur.execute(sql.encode()) await conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 if count > 0: @@ -269,7 +268,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (key,)) + cur.execute(sql.encode(), (key,)) row = cur.fetchone() if row is None: @@ -283,7 +282,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | SET expires_at = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ - conn.execute(as_query(update_sql), (new_expires_at, key)) + conn.execute(update_sql.encode(), (new_expires_at, key)) conn.commit() return bytes(row["data"]) @@ -303,21 +302,21 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No """ with self._config.provide_connection() as conn: - conn.execute(as_query(sql), (key, data, expires_at)) + conn.execute(sql.encode(), (key, data, expires_at)) conn.commit() def _delete(self, key: str) -> None: sql = f"DELETE FROM {self._table_name} WHERE session_id = %s" with self._config.provide_connection() as conn: - conn.execute(as_query(sql), (key,)) + conn.execute(sql.encode(), (key,)) conn.commit() def _delete_all(self) -> None: sql = f"DELETE FROM {self._table_name}" with self._config.provide_connection() as conn: - conn.execute(as_query(sql)) + conn.execute(sql.encode()) conn.commit() self._log_delete_all() @@ -329,7 +328,7 @@ def _exists(self, key: str) -> bool: """ with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(as_query(sql), (key,)) + cur.execute(sql.encode(), (key,)) row = cur.fetchone() return row is not None @@ -341,7 +340,7 @@ def _expires_in(self, key: str) -> "int | None": with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: - cur.execute(as_query(sql), (key,)) + cur.execute(sql.encode(), (key,)) row = cur.fetchone() if row is None or row["expires_at"] is None: @@ -360,7 +359,7 @@ def _delete_expired(self) -> int: sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP" with self._config.provide_connection() as conn, conn.cursor() as cur: - cur.execute(as_query(sql)) + cur.execute(sql.encode()) conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 if count > 0: diff --git a/sqlspec/adapters/psqlpy/_typing.py b/sqlspec/adapters/psqlpy/_typing.py index 31ca76935..3b04e90b0 100644 --- a/sqlspec/adapters/psqlpy/_typing.py +++ b/sqlspec/adapters/psqlpy/_typing.py @@ -16,22 +16,33 @@ class _PsqlpyUnavailableError(Exception): if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias - from psqlpy import Connection as PsqlpyConnection + from psqlpy import Connection as _PsqlpyConnection from psqlpy import ConnectionPool as PsqlpyConnectionPool - from psqlpy import Listener as PsqlpyListener - from psqlpy.exceptions import ConnectionExecuteError as PsqlpyConnectionExecuteError - from psqlpy.exceptions import DatabaseError as PsqlpyDatabaseError - from psqlpy.exceptions import DataError as PsqlpyDataError - from psqlpy.exceptions import Error as PsqlpyError - from psqlpy.exceptions import IntegrityError as PsqlpyIntegrityError - from psqlpy.exceptions import NotSupportedError as PsqlpyNotSupportedError - from psqlpy.exceptions import OperationalError as PsqlpyOperationalError + from psqlpy import Listener as _PsqlpyListener + from psqlpy.exceptions import ConnectionExecuteError as _PsqlpyConnectionExecuteError + from psqlpy.exceptions import DatabaseError as _PsqlpyDatabaseError + from psqlpy.exceptions import DataError as _PsqlpyDataError + from psqlpy.exceptions import Error as _PsqlpyError + from psqlpy.exceptions import IntegrityError as _PsqlpyIntegrityError + from psqlpy.exceptions import NotSupportedError as _PsqlpyNotSupportedError + from psqlpy.exceptions import OperationalError as _PsqlpyOperationalError from psqlpy.extra_types import JSONB as PSQLPY_JSONB from sqlspec.adapters.psqlpy.driver import PsqlpyDriver from sqlspec.core import StatementConfig + PsqlpyConnection: TypeAlias = _PsqlpyConnection + PsqlpyDataError: TypeAlias = _PsqlpyDataError + PsqlpyDatabaseError: TypeAlias = _PsqlpyDatabaseError + PsqlpyConnectionExecuteError: TypeAlias = _PsqlpyConnectionExecuteError + PsqlpyError: TypeAlias = _PsqlpyError + PsqlpyIntegrityError: TypeAlias = _PsqlpyIntegrityError + PsqlpyListener: TypeAlias = _PsqlpyListener + PsqlpyNotSupportedError: TypeAlias = _PsqlpyNotSupportedError + PsqlpyOperationalError: TypeAlias = _PsqlpyOperationalError + if not TYPE_CHECKING: PsqlpyConnection = import_optional_attr("psqlpy", "Connection") or Any diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index 1fe07aed4..f525a7c50 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -13,7 +13,6 @@ from sqlglot.errors import ParseError from sqlspec.adapters.psqlpy._typing import PsqlpyDataError, PsqlpyIntegrityError, PsqlpyOperationalError -from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector from sqlspec.core import ( DriverParameterProfile, ParameterStyle, @@ -92,7 +91,6 @@ "TIMESTAMP WITHOUT TIME ZONE", }) _UUID_CASTS: Final[frozenset[str]] = frozenset({"UUID"}) -_VECTOR_CASTS: Final[frozenset[str]] = frozenset({"VECTOR", "HALFVEC", "SPARSEVEC"}) _DECIMAL_NORMALIZER = build_nested_decimal_normalizer(mode="float") _JSONB_TYPE: type[Any] | None = None try: @@ -151,7 +149,7 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str connection_config: Raw connection configuration mapping. Returns: - Dictionary with sanitized connection parameters accepted by psqlpy. + Dictionary with connection parameters. """ config = {key: value for key, value in connection_config.items() if value is not None} dsn = ( @@ -208,11 +206,10 @@ def collect_rows(query_result: Any | None) -> "tuple[list[dict[str, Any]], list[ Returns: Tuple of (rows, column_names). """ - dict_rows: list[dict[str, Any]] = ( - cast("list[dict[str, Any]]", query_result if isinstance(query_result, list) else query_result.result()) - if query_result - else [] - ) + if not query_result: + return [], [] + + dict_rows = cast("list[dict[str, Any]]", query_result if isinstance(query_result, list) else query_result.result()) if not dict_rows: return [], [] return dict_rows, list(dict_rows[0]) @@ -227,13 +224,6 @@ class PsqlpyStreamSource: left untouched. """ - _chunk_size: int - _cursor: Any - _driver: Any - _parameters: Any - _sql: str - _transaction: Any - __slots__ = ("_chunk_size", "_cursor", "_driver", "_parameters", "_sql", "_transaction") def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: @@ -241,8 +231,8 @@ def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> N self._sql = sql self._parameters = parameters self._chunk_size = chunk_size - self._cursor = None - self._transaction = None + self._cursor: Any = None + self._transaction: Any = None async def start(self) -> None: handler = self._driver.handle_database_exceptions() @@ -296,7 +286,6 @@ async def close(self, error: bool = False) -> None: def coerce_numeric_for_write(value: Any) -> Any: - """Coerce numerical values to Decimal for precise Postgres numeric writes.""" if isinstance(value, float): return decimal.Decimal(str(value)) if isinstance(value, decimal.Decimal): @@ -609,8 +598,6 @@ def _coerce_parameter_for_cast(value: Any, cast_type: str, serializer: "Callable return _coerce_uuid_parameter(value) if upper_cast in _TIMESTAMP_CASTS: return _coerce_timestamp_parameter(value) - if upper_cast in _VECTOR_CASTS: - return coerce_pgvector(value) return value diff --git a/sqlspec/adapters/psqlpy/data_dictionary.py b/sqlspec/adapters/psqlpy/data_dictionary.py index a1b231d5d..05baa87d9 100644 --- a/sqlspec/adapters/psqlpy/data_dictionary.py +++ b/sqlspec/adapters/psqlpy/data_dictionary.py @@ -3,6 +3,8 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast +from mypy_extensions import mypyc_attr + from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -66,6 +68,7 @@ } +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsqlpyDataDictionary(AsyncDataDictionaryBase): """PostgreSQL-specific async data dictionary via psqlpy.""" diff --git a/sqlspec/adapters/psqlpy/driver.py b/sqlspec/adapters/psqlpy/driver.py index 1526f783a..69b9f88c0 100644 --- a/sqlspec/adapters/psqlpy/driver.py +++ b/sqlspec/adapters/psqlpy/driver.py @@ -4,7 +4,6 @@ and transaction management. """ -from collections.abc import Mapping from typing import TYPE_CHECKING, Any, cast from sqlspec.adapters.psqlpy._typing import PsqlpyCursor, PsqlpyDatabaseError, PsqlpyError, PsqlpySessionContext @@ -26,28 +25,13 @@ split_schema_and_table, ) from sqlspec.adapters.psqlpy.data_dictionary import PsqlpyDataDictionary -from sqlspec.core import ( - SQL, - StackResult, - StatementConfig, - StatementStack, - get_cache_config, - is_copy_operation, - register_driver_profile, -) -from sqlspec.driver import ( - AsyncDriverAdapterBase, - AsyncRowStream, - BaseAsyncExceptionHandler, - StackExecutionObserver, - describe_stack_statement, -) -from sqlspec.exceptions import SQLSpecError, StackExecutionError -from sqlspec.utils.logging import get_logger +from sqlspec.core import SQL, StatementConfig, get_cache_config, register_driver_profile +from sqlspec.driver import AsyncDriverAdapterBase, AsyncRowStream, BaseAsyncExceptionHandler +from sqlspec.exceptions import SQLSpecError from sqlspec.utils.text import normalize_identifier, quote_identifier if TYPE_CHECKING: - from collections.abc import Sequence + from collections.abc import Mapping, Sequence from sqlspec.adapters.psqlpy._typing import PsqlpyConnection from sqlspec.core import ArrowResult, SQLResult @@ -56,8 +40,6 @@ __all__ = ("PsqlpyCursor", "PsqlpyDriver", "PsqlpyExceptionHandler", "PsqlpySessionContext") -logger = get_logger("sqlspec.adapters.psqlpy") - class PsqlpyExceptionHandler(BaseAsyncExceptionHandler): """Async context manager for handling psqlpy database exceptions. @@ -89,11 +71,7 @@ class PsqlpyDriver(AsyncDriverAdapterBase): and transaction management. """ - _data_dictionary: PsqlpyDataDictionary | None - _transaction_active: bool - _json_columns_cache: dict[tuple[str | None, str], set[str]] - - __slots__ = ("_data_dictionary", "_json_columns_cache", "_transaction_active") + __slots__ = ("_data_dictionary", "_transaction_active") dialect = "postgres" def __init__( @@ -108,9 +86,8 @@ def __init__( ) super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features) - self._data_dictionary = None + self._data_dictionary: PsqlpyDataDictionary | None = None self._transaction_active = False - self._json_columns_cache = {} async def dispatch_execute(self, cursor: "PsqlpyConnection", statement: SQL) -> "ExecutionResult": """Execute single SQL statement. @@ -142,7 +119,14 @@ async def dispatch_execute(self, cursor: "PsqlpyConnection", statement: SQL) -> count_sql = _dml_count_query(sql) if count_sql is not None: count_result = await cursor.fetch(count_sql, params) - rows_affected = _extract_dml_count(count_result) + count_rows, _ = collect_rows(count_result) + if len(count_rows) != 1 or set(count_rows[0]) != {_DML_COUNT_COLUMN}: + msg = "psqlpy DML row count query returned an invalid result" + raise SQLSpecError(msg) + rows_affected = count_rows[0][_DML_COUNT_COLUMN] + if type(rows_affected) is not int or rows_affected < 0: + msg = "psqlpy DML row count query returned an invalid count" + raise SQLSpecError(msg) return self.create_execution_result(cursor, rowcount_override=rows_affected) result = await cursor.execute(sql, params) @@ -174,7 +158,7 @@ async def dispatch_execute_many(self, cursor: "PsqlpyConnection", statement: SQL return self.create_execution_result(cursor, rowcount_override=rows_affected, is_many_result=True) async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: SQL) -> "ExecutionResult": - """Execute SQL script with statement splitting and sequential execution. + """Execute SQL script with statement splitting. Args: cursor: Psqlpy connection object @@ -186,7 +170,6 @@ async def dispatch_execute_script(self, cursor: "PsqlpyConnection", statement: S sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) prepared_parameters = cast("Sequence[Any] | Mapping[str, Any] | None", prepared_parameters) statement_config = statement.statement_config - statements = self.split_script_statements(sql, statement_config, strip_trailing_semicolon=True) successful_count = 0 @@ -264,10 +247,6 @@ async def _resolve_json_columns(self, schema_name: "str | None", table_name: str Returns: Names of columns typed json or jsonb. """ - cache_key = (schema_name, table_name) - if cache_key in self._json_columns_cache: - return self._json_columns_cache[cache_key] - qualified = quote_identifier(table_name) if schema_name is not None: qualified = f"{quote_identifier(schema_name)}.{qualified}" @@ -282,9 +261,7 @@ async def _resolve_json_columns(self, schema_name: "str | None", table_name: str [qualified], ) data, _ = collect_rows(rows) - result = {str(row["column_name"]) for row in data} - self._json_columns_cache[cache_key] = result - return result + return {str(row["column_name"]) for row in data} async def has_schema(self, schema: str) -> bool: """Return whether a PostgreSQL schema exists.""" @@ -322,147 +299,6 @@ def handle_database_exceptions(self) -> "PsqlpyExceptionHandler": """ return PsqlpyExceptionHandler() - async def execute_stack( - self, stack: "StatementStack", *, continue_on_error: bool = False - ) -> "tuple[StackResult, ...]": - """Execute a StatementStack using psqlpy transaction pipeline when supported.""" - if ( - not isinstance(stack, StatementStack) - or not stack - or self.stack_native_disabled - or continue_on_error - or not hasattr(self.connection, "transaction") - ): - return await super().execute_stack(stack, continue_on_error=continue_on_error) - - prepared_ops = self._prepare_pipeline_operations(stack) - if prepared_ops is None: - return await super().execute_stack(stack, continue_on_error=continue_on_error) - - return await self._execute_stack_pipeline(stack, prepared_ops) - - def _prepare_pipeline_operations(self, stack: "StatementStack") -> "list[tuple[SQL, str, list[Any], bool]] | None": - """Prepare stack operations for native psqlpy transaction pipeline execution. - - Returns None when any operation in the stack requires the sequential - fallback path (non-execute methods, per-operation statement_config, - scripts, batch operations, COPY operations, or mapping parameters). - """ - prepared: list[tuple[SQL, str, list[Any], bool]] = [] - for operation in stack.operations: - if operation.method != "execute": - return None - if operation.keyword_arguments and "statement_config" in operation.keyword_arguments: - return None - - kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else None - sql_statement = self.prepare_statement( - operation.statement, operation.arguments, statement_config=self.statement_config, kwargs=kwargs - ) - if sql_statement.is_script or sql_statement.is_many or is_copy_operation(sql_statement.operation_type): - return None - - sql, prepared_parameters = self._compiled_sql(sql_statement, self.statement_config) - if isinstance(prepared_parameters, Mapping): - return None - - params = list(prepared_parameters) if isinstance(prepared_parameters, (list, tuple)) else [] - is_dml_count = False - if not sql_statement.returns_rows() and sql_statement.operation_type in {"INSERT", "UPDATE", "DELETE"}: - try: - count_sql = _dml_count_query(sql) - except SQLSpecError: - return None - if count_sql is not None: - sql = count_sql - is_dml_count = True - - prepared.append((sql_statement, sql, params, is_dml_count)) - return prepared - - async def _execute_stack_pipeline( - self, stack: "StatementStack", prepared_ops: "list[tuple[SQL, str, list[Any], bool]]" - ) -> "tuple[StackResult, ...]": - """Execute prepared stack operations through psqlpy's native Rust pipeline.""" - results: list[StackResult] = [] - started_transaction = False - queries: list[tuple[str, list[Any] | None]] = [(sql, params) for _, sql, params, _ in prepared_ops] - - with StackExecutionObserver(self, stack, continue_on_error=False, native_pipeline=True): - try: - if not self._connection_in_transaction(): - await self.begin() - started_transaction = True - - transaction = self.connection.transaction() - exc_handler = self.handle_database_exceptions() - try: - query_results = await self._run_with_exception_handler( - exc_handler, transaction.pipeline, queries, True - ) - self._check_pending_exception(exc_handler) - except Exception as exc: - stack_error = StackExecutionError( - 0, - describe_stack_statement(stack.operations[0].statement), - exc, - adapter=type(self).__name__, - mode="fail-fast", - native_pipeline=True, - ) - raise stack_error from exc - - assert query_results is not None - for index, ((sql_statement, _, _, is_dml_count), query_result) in enumerate( - zip(prepared_ops, query_results, strict=False) - ): - try: - if sql_statement.returns_rows(): - dict_rows, column_names = collect_rows(query_result) - execution_result = self.create_execution_result( - self.connection, - selected_data=dict_rows, - column_names=column_names, - data_row_count=len(dict_rows), - is_select_result=True, - row_format="dict", - ) - elif is_dml_count: - rows_affected = _extract_dml_count(query_result) - execution_result = self.create_execution_result( - self.connection, rowcount_override=rows_affected - ) - else: - rows_affected = extract_rows_affected(query_result) - execution_result = self.create_execution_result( - self.connection, rowcount_override=rows_affected - ) - except Exception as exc: - stack_error = StackExecutionError( - index, - describe_stack_statement(stack.operations[index].statement), - exc, - adapter=type(self).__name__, - mode="fail-fast", - native_pipeline=True, - ) - raise stack_error from exc - - sql_result = self.build_statement_result(sql_statement, execution_result) - results.append(StackResult.from_sql_result(sql_result)) - - if started_transaction: - await self.commit() - except Exception: - if started_transaction: - try: - await self.rollback() - except Exception as rollback_error: - logger.debug("Rollback after psqlpy pipeline failure failed: %s", rollback_error) - raise - - return tuple(results) - async def select_to_storage( self, statement: "SQL | str", @@ -628,17 +464,4 @@ def _connection_in_transaction(self) -> bool: return self._transaction_active -def _extract_dml_count(count_result: Any) -> int: - """Validate and extract the affected row count from a psqlpy DML count CTE query result.""" - count_rows, _ = collect_rows(count_result) - if len(count_rows) != 1 or set(count_rows[0]) != {_DML_COUNT_COLUMN}: - msg = "psqlpy DML row count query returned an invalid result" - raise SQLSpecError(msg) - rows_affected = count_rows[0][_DML_COUNT_COLUMN] - if type(rows_affected) is not int or rows_affected < 0: - msg = "psqlpy DML row count query returned an invalid count" - raise SQLSpecError(msg) - return rows_affected - - register_driver_profile("psqlpy", driver_profile) diff --git a/sqlspec/adapters/psqlpy/litestar/store.py b/sqlspec/adapters/psqlpy/litestar/store.py index 664307112..ff40e16f7 100644 --- a/sqlspec/adapters/psqlpy/litestar/store.py +++ b/sqlspec/adapters/psqlpy/litestar/store.py @@ -82,8 +82,8 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by Returns: Session data as bytes if found and not expired, None otherwise. """ - if renew_for is not None and self._calculate_expires_at(renew_for) is not None: - new_expires_at = self._calculate_expires_at(renew_for) + new_expires_at = self._calculate_expires_at(renew_for) if renew_for is not None else None + if new_expires_at is not None: sql = f""" UPDATE {self._table_name} SET expires_at = CASE WHEN expires_at IS NOT NULL THEN $1 ELSE expires_at END, diff --git a/sqlspec/adapters/psqlpy/type_converter.py b/sqlspec/adapters/psqlpy/type_converter.py index 5078bb1a8..5f788c0d9 100644 --- a/sqlspec/adapters/psqlpy/type_converter.py +++ b/sqlspec/adapters/psqlpy/type_converter.py @@ -1,37 +1,26 @@ -"""PostgreSQL-specific helpers for the psqlpy adapter.""" +"""PostgreSQL-specific helpers for the psqlpy adapter. -from typing import TYPE_CHECKING, Any +This module preserves the ``register_pgvector`` placeholder used by the +driver configuration layer. +""" -from sqlspec.typing import PGVECTOR_INSTALLED, import_optional_attr +from typing import TYPE_CHECKING + +from sqlspec.typing import PGVECTOR_INSTALLED if TYPE_CHECKING: from sqlspec.adapters.psqlpy._typing import PsqlpyConnection as Connection -__all__ = ("coerce_pgvector", "register_pgvector") - - -def coerce_pgvector(value: Any) -> Any: - """Coerce sequence or numpy array to psqlpy PgVector.""" - if value is None or not PGVECTOR_INSTALLED: - return value - pg_vector_cls = import_optional_attr("psqlpy.extra_types", "PgVector") - if pg_vector_cls is None: - return value - try: - if isinstance(value, pg_vector_cls): - return value - if isinstance(value, (list, tuple)): - return pg_vector_cls(list(value)) - if hasattr(value, "tolist"): - return pg_vector_cls(value.tolist()) - except Exception: - return value - return value +__all__ = ("register_pgvector",) def register_pgvector(connection: "Connection") -> None: """Register pgvector type handlers on psqlpy connection. + Currently a placeholder for future implementation. The psqlpy library + does not yet expose a type handler registration API compatible with + pgvector's automatic conversion system. + Args: connection: Psqlpy connection instance. """ diff --git a/sqlspec/adapters/psycopg/_typing.py b/sqlspec/adapters/psycopg/_typing.py index 323c6b299..f59eb308c 100644 --- a/sqlspec/adapters/psycopg/_typing.py +++ b/sqlspec/adapters/psycopg/_typing.py @@ -5,7 +5,7 @@ """ import contextlib -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol import psycopg as psycopg_module from psycopg import AsyncConnection, AsyncCursor, Connection, Cursor @@ -26,9 +26,7 @@ from psycopg.sql import Identifier as PsycopgIdentifier from psycopg.types.json import Jsonb as PsycopgJsonb from psycopg_pool import AsyncConnectionPool as PsycopgAsyncConnectionPool -from psycopg_pool import AsyncNullConnectionPool as PsycopgAsyncNullConnectionPool from psycopg_pool import ConnectionPool as PsycopgConnectionPool -from psycopg_pool import NullConnectionPool as PsycopgNullConnectionPool from psycopg_pool.abc import AsyncConnectFailedCB as PsycopgAsyncConnectFailedCB from psycopg_pool.abc import AsyncConnectionCB as PsycopgAsyncConnectionCB from psycopg_pool.abc import ConnectFailedCB as PsycopgConnectFailedCB @@ -44,7 +42,8 @@ from google.cloud.alloydb.connector import Connector as PsycopgAlloydbConnector from sqlspec.adapters.psycopg.driver import PsycopgAsyncDriver, PsycopgSyncDriver - from sqlspec.core import StatementConfig + from sqlspec.builder import QueryBuilder + from sqlspec.core import SQL, Statement, StatementConfig PsycopgSyncConnection: TypeAlias = Connection[PsycopgDictRow] PsycopgAsyncConnection: TypeAlias = AsyncConnection[PsycopgDictRow] @@ -66,7 +65,6 @@ "PsycopgAsyncConnectionCB", "PsycopgAsyncConnectionPool", "PsycopgAsyncCursor", - "PsycopgAsyncNullConnectionPool", "PsycopgAsyncRawCursor", "PsycopgAsyncRowFactory", "PsycopgAsyncSessionContext", @@ -81,7 +79,7 @@ "PsycopgJsonb", "PsycopgNativeAsyncConnection", "PsycopgNativeAsyncCursor", - "PsycopgNullConnectionPool", + "PsycopgPipelineDriver", "PsycopgProgrammingError", "PsycopgRowFactory", "PsycopgSQL", @@ -137,6 +135,31 @@ async def __aexit__( await self.cursor.close() +class PsycopgPipelineDriver(Protocol): + """Protocol for psycopg pipeline driver methods used in stack execution.""" + + statement_config: "StatementConfig" + + def prepare_statement( + self, + statement: "SQL | Statement | QueryBuilder", + parameters: Any, + *, + statement_config: "StatementConfig | None" = None, + kwargs: "dict[str, Any] | None" = None, + ) -> "SQL": ... + + def prepare_driver_parameters( + self, + parameters: Any, + statement_config: "StatementConfig", + is_many: bool = False, + prepared_statement: Any | None = None, + ) -> Any: ... + + def _compiled_sql(self, statement: "SQL", statement_config: "StatementConfig") -> "tuple[str, Any]": ... + + class PsycopgSyncSessionContext: """Sync context manager for psycopg sessions. diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index d493b2480..c8395c7e9 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -4,8 +4,6 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypedDict, cast from mypy_extensions import mypyc_attr -from psycopg.adapt import AdaptersMap -from psycopg.types.json import set_json_dumps, set_json_loads from typing_extensions import NotRequired, Self from sqlspec.adapters.psycopg._typing import ( @@ -18,9 +16,7 @@ PsycopgSyncSessionContext, ) from sqlspec.adapters.psycopg._typing import PsycopgAsyncConnectionPool as AsyncConnectionPool -from sqlspec.adapters.psycopg._typing import PsycopgAsyncNullConnectionPool as AsyncNullConnectionPool from sqlspec.adapters.psycopg._typing import PsycopgConnectionPool as ConnectionPool -from sqlspec.adapters.psycopg._typing import PsycopgNullConnectionPool as NullConnectionPool from sqlspec.adapters.psycopg.core import apply_driver_features, default_statement_config from sqlspec.adapters.psycopg.driver import ( PsycopgAsyncDriver, @@ -47,7 +43,6 @@ from sqlspec.extensions.events import EventRuntimeHints from sqlspec.typing import ALLOYDB_CONNECTOR_INSTALLED from sqlspec.utils.config_tools import normalize_connection_config -from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Awaitable, Callable, Mapping @@ -144,7 +139,6 @@ class PsycopgPoolParams(PsycopgConnectionParams): close_returns: NotRequired[bool] reconnect_failed: NotRequired["ConnectFailedCB | AsyncConnectFailedCB | None"] kwargs: NotRequired["dict[str, Any]"] - null_pool: NotRequired[bool] class PsycopgDriverFeatures(TypedDict): @@ -189,7 +183,6 @@ class PsycopgDriverFeatures(TypedDict): enable_alloydb_iam_auth: Enable AlloyDB IAM database authentication for sync connector connections. Defaults to False. alloydb_ip_type: AlloyDB connector IP type. Defaults to PRIVATE. - null_pool: Enable NullConnectionPool / AsyncNullConnectionPool for serverless / PgBouncer environments. """ enable_pgvector: NotRequired[bool] @@ -204,7 +197,6 @@ class PsycopgDriverFeatures(TypedDict): alloydb_instance_uri: NotRequired[str] enable_alloydb_iam_auth: NotRequired[bool] alloydb_ip_type: NotRequired[str] - null_pool: NotRequired[bool] def build_connection_config(connection_config: "PsycopgPoolParams | Mapping[str, Any] | None") -> dict[str, Any]: @@ -430,19 +422,10 @@ def _create_pool(self) -> "ConnectionPool": self._setup_alloydb_connector(all_config, pool_parameters) conninfo = None - is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) - pool_cls = NullConnectionPool if is_null_pool else ConnectionPool - if is_null_pool: - pool_parameters.pop("min_size", None) - pool_parameters.pop("max_size", None) - pool_parameters.pop("max_idle", None) - pool_parameters.pop("max_waiting", None) - pool_parameters.pop("num_workers", None) - if conninfo: - pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) + pool = ConnectionPool(conninfo, kwargs=all_config, **pool_parameters) else: - pool = pool_cls("", kwargs=all_config, **pool_parameters) + pool = ConnectionPool("", kwargs=all_config, **pool_parameters) return pool @@ -451,13 +434,7 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting - serializer = self.driver_features.get("json_serializer", to_json) - deserializer = self.driver_features.get("json_deserializer", from_json) - if isinstance(getattr(conn, "adapters", None), AdaptersMap): - with suppress(Exception): - set_json_dumps(serializer, conn) - set_json_loads(deserializer, conn) - + # Detect extensions on first connection, update dialect if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -761,20 +738,10 @@ async def _create_pool(self) -> "AsyncConnectionPool": conninfo = all_config.pop("conninfo", None) kwargs = all_config.pop("kwargs", {}) all_config.update(kwargs) - - is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) - pool_cls = AsyncNullConnectionPool if is_null_pool else AsyncConnectionPool - if is_null_pool: - pool_parameters.pop("min_size", None) - pool_parameters.pop("max_size", None) - pool_parameters.pop("max_idle", None) - pool_parameters.pop("max_waiting", None) - pool_parameters.pop("num_workers", None) - if conninfo: - pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) + pool = AsyncConnectionPool(conninfo, kwargs=all_config, **pool_parameters) else: - pool = pool_cls("", kwargs=all_config, **pool_parameters) + pool = AsyncConnectionPool("", kwargs=all_config, **pool_parameters) if open_pool is True: await pool.open() @@ -786,13 +753,7 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) - serializer = self.driver_features.get("json_serializer", to_json) - deserializer = self.driver_features.get("json_deserializer", from_json) - if isinstance(getattr(conn, "adapters", None), AdaptersMap): - with suppress(Exception): - set_json_dumps(serializer, conn) - set_json_loads(deserializer, conn) - + # Detect extensions on first connection, update dialect if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) diff --git a/sqlspec/adapters/psycopg/core.py b/sqlspec/adapters/psycopg/core.py index 13f806b92..9f4d35290 100644 --- a/sqlspec/adapters/psycopg/core.py +++ b/sqlspec/adapters/psycopg/core.py @@ -80,21 +80,11 @@ "resolve_runtime_statement_config", ) -TRANSACTION_STATUS_IDLE: int = 0 -TRANSACTION_STATUS_ACTIVE: int = 1 -TRANSACTION_STATUS_INTRANS: int = 2 -TRANSACTION_STATUS_INERROR: int = 3 -TRANSACTION_STATUS_UNKNOWN: int = 4 -try: - from psycopg.pq import TransactionStatus - - TRANSACTION_STATUS_IDLE = int(TransactionStatus.IDLE) - TRANSACTION_STATUS_ACTIVE = int(TransactionStatus.ACTIVE) - TRANSACTION_STATUS_INTRANS = int(TransactionStatus.INTRANS) - TRANSACTION_STATUS_INERROR = int(TransactionStatus.INERROR) - TRANSACTION_STATUS_UNKNOWN = int(TransactionStatus.UNKNOWN) -except (ImportError, AttributeError): - pass +TRANSACTION_STATUS_IDLE = 0 +TRANSACTION_STATUS_ACTIVE = 1 +TRANSACTION_STATUS_INTRANS = 2 +TRANSACTION_STATUS_INERROR = 3 +TRANSACTION_STATUS_UNKNOWN = 4 class PreparedStackOperation(NamedTuple): @@ -129,7 +119,6 @@ def pipeline_supported() -> bool: def build_copy_from_command(table: str, columns: "list[str]") -> "PsycopgComposed": - """Build a COPY FROM STDIN command.""" table_identifier = _compose_table_identifier(table) column_sql = PsycopgSQL(", ").join([PsycopgIdentifier(column) for column in columns]) return PsycopgSQL("COPY {} ({}) FROM STDIN").format(table_identifier, column_sql) diff --git a/sqlspec/adapters/psycopg/data_dictionary.py b/sqlspec/adapters/psycopg/data_dictionary.py index 47522ce58..80c16accc 100644 --- a/sqlspec/adapters/psycopg/data_dictionary.py +++ b/sqlspec/adapters/psycopg/data_dictionary.py @@ -3,6 +3,8 @@ from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast +from mypy_extensions import mypyc_attr + from sqlspec.data_dictionary import ( ColumnMetadata, DDLResult, @@ -66,6 +68,7 @@ } +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsycopgSyncDataDictionary(SyncDataDictionaryBase): """PostgreSQL-specific sync data dictionary.""" @@ -342,6 +345,7 @@ def get_foreign_keys( ) +@mypyc_attr(allow_interpreted_subclasses=True, native_class=False) class PsycopgAsyncDataDictionary(AsyncDataDictionaryBase): """PostgreSQL-specific async data dictionary.""" diff --git a/sqlspec/adapters/psycopg/driver.py b/sqlspec/adapters/psycopg/driver.py index a52aab061..7f48b38f2 100644 --- a/sqlspec/adapters/psycopg/driver.py +++ b/sqlspec/adapters/psycopg/driver.py @@ -67,37 +67,12 @@ if TYPE_CHECKING: from collections import abc - from typing import Protocol - from sqlspec.builder import QueryBuilder - from sqlspec.core import ArrowResult, Statement + from sqlspec.adapters.psycopg._typing import PsycopgPipelineDriver + from sqlspec.core import ArrowResult from sqlspec.driver import CachedQuery, ExecutionResult from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry - class PsycopgPipelineDriver(Protocol): - """Protocol for psycopg pipeline driver methods used in stack execution.""" - - statement_config: "StatementConfig" - - def prepare_statement( - self, - statement: "SQL | Statement | QueryBuilder", - parameters: Any, - *, - statement_config: "StatementConfig | None" = None, - kwargs: "dict[str, Any] | None" = None, - ) -> "SQL": ... - - def prepare_driver_parameters( - self, - parameters: Any, - statement_config: "StatementConfig", - is_many: bool = False, - prepared_statement: Any | None = None, - ) -> Any: ... - - def _compiled_sql(self, statement: "SQL", statement_config: "StatementConfig") -> "tuple[str, Any]": ... - __all__ = ( "PsycopgAsyncCursor", diff --git a/tests/integration/adapters/_shared/_cases.py b/tests/integration/adapters/_shared/_cases.py index b8e269d2c..b1419ca52 100644 --- a/tests/integration/adapters/_shared/_cases.py +++ b/tests/integration/adapters/_shared/_cases.py @@ -744,11 +744,13 @@ def _db2_case(mode: Literal["sync", "async"], marks: tuple[Mark | MarkDecorator, supports_connection_hook=True, config_factory_fixture="lifecycle_config_asyncpg", supports_connection_instance=True, + native_stack_parity_mode="standard", extra_assertions=( "explain_modifiers:postgres", "arrow_specifics:postgres", "execute_many_specifics:postgres", "param_codecs:asyncpg", + "statement_stack:native_fallback_parity", "streaming_native:asyncpg", "stream_error_close:pg", ), diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 66bb8a7ca..0f3480561 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -219,7 +219,6 @@ class SourceEquivalenceCase: "alloydb_instance_uri", "enable_alloydb_iam_auth", "alloydb_ip_type", - "null_pool", ), "pymssql": ("json_serializer", "json_deserializer", "on_connection_create", "enable_events", "events_backend"), "pymysql": ( diff --git a/tests/integration/adapters/postgres/asyncpg/test_driver.py b/tests/integration/adapters/postgres/asyncpg/test_driver.py index b303bbcb6..31239e022 100644 --- a/tests/integration/adapters/postgres/asyncpg/test_driver.py +++ b/tests/integration/adapters/postgres/asyncpg/test_driver.py @@ -510,6 +510,25 @@ async def test_asyncpg_statement_stack_continue_on_error_inside_transaction(asyn assert [row["id"] for row in persisted.get_data()] == [1, 2] +async def test_asyncpg_statement_stack_marks_prepared(asyncpg_session: "AsyncpgDriver") -> None: + """Prepared statement metadata should be attached to stack results.""" + + await asyncpg_session.execute_script("DELETE FROM test_table_asyncpg") + + stack = ( + StatementStack() + .push_execute("INSERT INTO test_table_asyncpg (id, name, value) VALUES ($1, $2, $3)", (1, "stack-prepared", 50)) + .push_execute("SELECT value FROM test_table_asyncpg WHERE id = $1", (1,)) + ) + + results = await asyncpg_session.execute_stack(stack) + + assert results[0].metadata is not None + assert results[0].metadata.get("prepared_statement") is True + assert results[1].metadata is not None + assert results[1].metadata.get("prepared_statement") is True + + async def test_asyncpg_pool_concurrency(postgres_service: PostgresService) -> None: """Verify that multiple concurrent calls to provide_pool result in a single pool.""" config_params = AsyncpgPoolConfig( diff --git a/tests/unit/adapters/test_asyncpg/test_config.py b/tests/unit/adapters/test_asyncpg/test_config.py index 76e4e189b..80bdca140 100644 --- a/tests/unit/adapters/test_asyncpg/test_config.py +++ b/tests/unit/adapters/test_asyncpg/test_config.py @@ -311,10 +311,3 @@ def test_asyncpg_config_normalizes_aliases() -> None: assert "conninfo" not in config.connection_config assert "dbname" not in config.connection_config assert "username" not in config.connection_config - - -def test_asyncpg_pgbouncer_connection_config_sets_statement_cache_size_zero() -> None: - """Enabling pgbouncer in connection_config should pop pgbouncer and disable statement caching.""" - config = AsyncpgConfig(connection_config={"dsn": "postgresql://localhost:5432/test", "pgbouncer": True}) - assert "pgbouncer" not in config.connection_config - assert config.connection_config["statement_cache_size"] == 0 diff --git a/tests/unit/adapters/test_contract_statement_stack_parity.py b/tests/unit/adapters/test_contract_statement_stack_parity.py index f8d21c94e..9359afd42 100644 --- a/tests/unit/adapters/test_contract_statement_stack_parity.py +++ b/tests/unit/adapters/test_contract_statement_stack_parity.py @@ -7,7 +7,7 @@ from tests.integration.adapters._shared.behaviors import STATEMENT_STACK_SCOPE, validate_extra_assertions PARITY_PROOF_KEY = "statement_stack:native_fallback_parity" -OPTED_IN_CASE_IDS = ("psycopg-sync", "psycopg-async", "oracledb-async") +OPTED_IN_CASE_IDS = ("psycopg-sync", "asyncpg-async", "psycopg-async", "oracledb-async") def test_sync_statement_stack_parity_proof_registered() -> None: diff --git a/tests/unit/adapters/test_psqlpy/test_transaction_state.py b/tests/unit/adapters/test_psqlpy/test_transaction_state.py index b92ea9cdd..9c48ea87f 100644 --- a/tests/unit/adapters/test_psqlpy/test_transaction_state.py +++ b/tests/unit/adapters/test_psqlpy/test_transaction_state.py @@ -7,10 +7,9 @@ from sqlspec.adapters.psqlpy._typing import PsqlpyDatabaseError from sqlspec.adapters.psqlpy.config import PsqlpyConfig -from sqlspec.adapters.psqlpy.core import _DML_COUNT_COLUMN, PsqlpyStreamSource +from sqlspec.adapters.psqlpy.core import PsqlpyStreamSource from sqlspec.adapters.psqlpy.driver import PsqlpyDriver -from sqlspec.core import StatementStack -from sqlspec.exceptions import SQLSpecError, StackExecutionError +from sqlspec.exceptions import SQLSpecError pytestmark = pytest.mark.anyio @@ -206,116 +205,3 @@ async def test_load_from_arrow_decodes_json_text_for_json_columns() -> None: _table_name, records, _kwargs = connection.copy_calls[0] assert records == [(1, {"name": "alpha"}, '{"not": "json"}')] - - -class _PipelineTransaction: - def __init__( - self, - pipeline_calls: "list[tuple[list[tuple[str, list[Any] | None]], bool]]", - results: "list[Any]", - error: "Exception | None" = None, - ) -> None: - self._pipeline_calls = pipeline_calls - self._results = results - self._error = error - - async def pipeline(self, queries: "list[tuple[str, list[Any] | None]]", prepared: bool = True) -> "list[Any]": - self._pipeline_calls.append((queries, prepared)) - if self._error is not None: - raise self._error - return self._results - - -class _PipelineConnection(_FakeConnection): - def __init__(self, results: "list[Any] | None" = None, error: "Exception | None" = None) -> None: - super().__init__() - self.pipeline_calls: list[tuple[list[tuple[str, list[Any] | None]], bool]] = [] - self.fetch_calls: list[tuple[str, Any]] = [] - self.execute_many_calls: list[tuple[str, Any]] = [] - self._results = results or [] - self._error = error - - def transaction(self) -> _PipelineTransaction: - return _PipelineTransaction(self.pipeline_calls, self._results, self._error) - - async def fetch(self, sql: str, parameters: Any = None) -> Any: - self.fetch_calls.append((sql, parameters)) - if _DML_COUNT_COLUMN in sql: - return SimpleNamespace(result=lambda: [{_DML_COUNT_COLUMN: 1}]) - return SimpleNamespace(result=lambda: [{"id": 1}]) - - async def execute_many(self, sql: str, parameters: Any) -> None: - self.execute_many_calls.append((sql, parameters)) - - -async def test_execute_stack_uses_native_transaction_pipeline() -> None: - """Supported execute stacks should run through connection.transaction().pipeline.""" - connection = _PipelineConnection( - results=[ - SimpleNamespace(result=lambda: [{_DML_COUNT_COLUMN: 2}]), - SimpleNamespace(result=lambda: [{"id": 1, "name": "alpha"}]), - ] - ) - driver = PsqlpyDriver(cast("Any", connection)) - stack = ( - StatementStack() - .push_execute("INSERT INTO items (name) VALUES ($1)", "alpha") - .push_execute("SELECT id, name FROM items WHERE name = $1", "alpha") - ) - - results = await driver.execute_stack(stack) - - assert len(results) == 2 - assert results[0].rows_affected == 2 - assert results[1].result is not None - assert results[1].result.get_data() == [{"id": 1, "name": "alpha"}] - assert len(connection.pipeline_calls) == 1 - queries, prepared = connection.pipeline_calls[0] - assert prepared is True - assert len(queries) == 2 - assert _DML_COUNT_COLUMN in queries[0][0] - assert queries[0][1] == ["alpha"] - assert queries[1] == ("SELECT id, name FROM items WHERE name = $1", ["alpha"]) - assert connection.statements == ["BEGIN", "COMMIT"] - - -async def test_execute_stack_falls_back_when_continue_on_error_or_non_execute() -> None: - """Stacks with continue_on_error or non-execute operations must fall back to sequential execution.""" - connection = _PipelineConnection() - driver = PsqlpyDriver(cast("Any", connection)) - - continue_stack = StatementStack().push_execute("SELECT 1") - await driver.execute_stack(continue_stack, continue_on_error=True) - assert connection.pipeline_calls == [] - assert len(connection.fetch_calls) == 1 - - many_stack = StatementStack().push_execute_many("INSERT INTO items (name) VALUES ($1)", [("a",), ("b",)]) - await driver.execute_stack(many_stack) - assert connection.pipeline_calls == [] - assert len(connection.execute_many_calls) == 1 - - -async def test_execute_stack_falls_back_when_native_stack_disabled() -> None: - """Native stack disablement must bypass connection.transaction().pipeline.""" - connection = _PipelineConnection() - driver = PsqlpyDriver(cast("Any", connection), driver_features={"stack_native_disabled": True}) - stack = StatementStack().push_execute("SELECT 1") - - await driver.execute_stack(stack) - - assert connection.pipeline_calls == [] - assert len(connection.fetch_calls) == 1 - - -async def test_execute_stack_native_pipeline_error_rolls_back_and_wraps() -> None: - """Pipeline failures should roll back owned transactions and raise StackExecutionError.""" - connection = _PipelineConnection(error=PsqlpyDatabaseError("unique constraint violation")) - driver = PsqlpyDriver(cast("Any", connection)) - stack = StatementStack().push_execute("INSERT INTO items (name) VALUES ($1)", "dup") - - with pytest.raises(StackExecutionError) as exc_info: - await driver.execute_stack(stack) - - assert exc_info.value.native_pipeline is True - assert connection.statements == ["BEGIN", "ROLLBACK"] - assert driver._connection_in_transaction() is False From 94a69c918eb9ab518e47aba6900d1aa3f8aea1c6 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:32:58 +0000 Subject: [PATCH 08/10] fix: preserve native PostgreSQL adapter controls --- docs/changelog.rst | 18 ++- docs/reference/adapters/asyncpg.rst | 19 +++ docs/reference/adapters/cockroach_asyncpg.rst | 8 ++ docs/reference/adapters/cockroach_psycopg.rst | 8 ++ docs/reference/adapters/psqlpy.rst | 9 ++ docs/reference/adapters/psycopg.rst | 13 ++ sqlspec/adapters/asyncpg/config.py | 28 +++- sqlspec/adapters/asyncpg/core.py | 38 ++++-- sqlspec/adapters/asyncpg/driver.py | 38 ++++-- sqlspec/adapters/cockroach_asyncpg/config.py | 6 +- sqlspec/adapters/cockroach_asyncpg/core.py | 27 ++++ sqlspec/adapters/cockroach_psycopg/config.py | 35 +---- sqlspec/adapters/cockroach_psycopg/core.py | 46 +++++++ sqlspec/adapters/psqlpy/core.py | 3 + sqlspec/adapters/psqlpy/type_converter.py | 20 ++- sqlspec/adapters/psycopg/_typing.py | 4 + sqlspec/adapters/psycopg/config.py | 44 +++++- .../adapters/test_postgres_native_options.py | 127 ++++++++++++++++++ 18 files changed, 419 insertions(+), 72 deletions(-) create mode 100644 tests/unit/adapters/test_postgres_native_options.py diff --git a/docs/changelog.rst b/docs/changelog.rst index 3174eff84..a5281a602 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,13 +1,13 @@ -========= +== Changelog -========= +== All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates -============== +======= Unreleased ---------- @@ -24,6 +24,13 @@ Unreleased * ADBC FlightSQL adds options for TLS/mTLS, RPC timeouts, message size, cookies and headers. Values set in native ``db_kwargs`` take precedence. +* 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. + * 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, @@ -38,6 +45,8 @@ Unreleased **Fixed:** +* Builder upserts emit ``MERGE`` for the ``db2`` dialect. + * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve question marks in quoted identifiers, literals, and comments. @@ -45,7 +54,6 @@ Unreleased parsing SQL again. ADBC keeps bound values in its ADK store queries. DuckDB Arrow loads keep sparse dictionary fields and quote table names. -* Builder upserts emit ``MERGE`` for the ``db2`` dialect. * 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. @@ -2229,7 +2237,7 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions -================= +========== For releases before ``v0.24.0``, see the repository tag history and GitHub release records. diff --git a/docs/reference/adapters/asyncpg.rst b/docs/reference/adapters/asyncpg.rst index e358841b7..9d45c76c9 100644 --- a/docs/reference/adapters/asyncpg.rst +++ b/docs/reference/adapters/asyncpg.rst @@ -89,3 +89,22 @@ All PostgreSQL adapters support ``data ? 'key'``, ``data ? $1``, ``data ? :key``, ``?|``, and ``?&`` without treating the operator as a placeholder. Write ``data ?? other_col`` for identifier or function right-hand operands. Write parameterized intervals as ``? * interval '1 day'`` (or use ``$1``). + +Native execution options +------------------------ + +Set ``execution_args={"timeout": seconds}`` on ``StatementConfig`` to forward +an asyncpg timeout to queries, batches, script statements, stack operations, +and stream fetches. ``command_timeout`` is accepted as an alias; ``timeout`` +takes precedence. Explicit ``0`` and ``None`` values are preserved. + +``driver_features={"pgbouncer": True}`` disables asyncpg's statement cache and +SQLSpec's explicit prepared statements in stacks, while preserving transaction +cleanup. The option is also accepted in ``connection_config``. Use this mode +when the proxy configuration does not support prepared statements; PgBouncer +can support them when configured to track protocol-level prepared statements. + +``driver_features["type_codecs"]`` accepts a list of native codec specifications. +Each entry requires ``typename``, ``encoder``, and ``decoder``; ``schema`` defaults +to ``public`` and ``format`` to ``text``. Codecs register after SQLSpec's built-in +JSON and vector setup and before ``on_connection_create``. diff --git a/docs/reference/adapters/cockroach_asyncpg.rst b/docs/reference/adapters/cockroach_asyncpg.rst index 629d5416b..48a3a3109 100644 --- a/docs/reference/adapters/cockroach_asyncpg.rst +++ b/docs/reference/adapters/cockroach_asyncpg.rst @@ -130,3 +130,11 @@ namespace: ``"litestar"``, ``"events"``, or ``"adk"`` as supported by this adapt .. autoclass:: sqlspec.adapters.cockroach_asyncpg.adk.CockroachAsyncpgADKConfig :members: :show-inheritance: + +Native startup settings +----------------------- + +``connection_config`` accepts ``application_name``, +``default_transaction_use_follower_reads`` (boolean), and ``results_buffer_size`` +(non-negative integer bytes). These map to asyncpg ``server_settings``; +explicit entries in ``server_settings`` take precedence. diff --git a/docs/reference/adapters/cockroach_psycopg.rst b/docs/reference/adapters/cockroach_psycopg.rst index db3eaf459..569232166 100644 --- a/docs/reference/adapters/cockroach_psycopg.rst +++ b/docs/reference/adapters/cockroach_psycopg.rst @@ -254,3 +254,11 @@ namespace: ``"litestar"``, ``"events"``, or ``"adk"`` as supported by this adapt .. autoclass:: sqlspec.adapters.cockroach_psycopg.adk.CockroachPsycopgADKConfig :members: :show-inheritance: + +Native startup settings +----------------------- + +``connection_config`` accepts ``default_transaction_use_follower_reads`` as a +boolean, plus non-negative integer ``results_buffer_size`` (bytes), +``statement_timeout`` and ``idle_in_transaction_session_timeout`` (milliseconds). +These settings append to the existing libpq ``options`` string. diff --git a/docs/reference/adapters/psqlpy.rst b/docs/reference/adapters/psqlpy.rst index 97a39c5b6..a7cdf1d9b 100644 --- a/docs/reference/adapters/psqlpy.rst +++ b/docs/reference/adapters/psqlpy.rst @@ -77,3 +77,12 @@ namespace: ``"litestar"``, ``"events"``, or ``"adk"`` as supported by this adapt .. autoclass:: sqlspec.adapters.psqlpy.adk.PsqlpyADKConfig :members: :show-inheritance: + +Dense vector parameters +----------------------- + +Parameters explicitly cast to ``VECTOR`` accept lists, tuples, and objects +with ``tolist()``. SQLSpec wraps these values with psqlpy's native ``PgVector``; +already-wrapped values pass through. This conversion does not require the +Python ``pgvector`` package and does not apply the dense-vector encoder to +``HALFVEC`` or ``SPARSEVEC``. diff --git a/docs/reference/adapters/psycopg.rst b/docs/reference/adapters/psycopg.rst index 405b2c0e8..f9e872a24 100644 --- a/docs/reference/adapters/psycopg.rst +++ b/docs/reference/adapters/psycopg.rst @@ -110,3 +110,16 @@ namespace: ``"litestar"``, ``"events"``, or ``"adk"`` as supported by this adapt .. autoclass:: sqlspec.adapters.psycopg.adk.PsycopgADKConfig :members: :show-inheritance: + +Native null pools and JSON codecs +-------------------------------- + +Set ``null_pool=True`` in ``connection_config`` or ``driver_features`` to select +psycopg's native null pool. Returned connections close instead of remaining idle. +``max_size`` and ``max_waiting`` still control concurrency and waiting clients; +``min_size`` is omitted because null pools do not maintain idle connections. +The same option works for sync and async configurations. + +``json_serializer`` and ``json_deserializer`` driver features also configure +psycopg's native JSON adaptation for each new connection. Codec setup errors +propagate to the caller. diff --git a/sqlspec/adapters/asyncpg/config.py b/sqlspec/adapters/asyncpg/config.py index f92543131..40443d24a 100644 --- a/sqlspec/adapters/asyncpg/config.py +++ b/sqlspec/adapters/asyncpg/config.py @@ -101,6 +101,7 @@ class AsyncpgConnectionConfig(TypedDict): connect_timeout: NotRequired[float] command_timeout: NotRequired[float] statement_cache_size: NotRequired[int] + pgbouncer: NotRequired[bool] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] server_settings: NotRequired["dict[str, str]"] @@ -191,6 +192,9 @@ class AsyncpgDriverFeatures(TypedDict): - "notify_queue": Durable queue plus a PostgreSQL notification wakeup hint - "poll_queue": Durable queue discovered by polling Defaults to "notify". + pgbouncer: Enable PgBouncer transaction-pooling compatibility mode. + Disables server-side prepared statement caching (statement_cache_size=0). + type_codecs: Optional list of custom type codec specifications to register. """ json_serializer: NotRequired["Callable[[Any], str]"] @@ -211,6 +215,8 @@ class AsyncpgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["notify", "notify_queue", "poll_queue"]] connection_instance: NotRequired["AsyncpgPool"] on_connection_create: NotRequired["Callable[[AsyncpgConnection], Awaitable[None]]"] + pgbouncer: NotRequired[bool] + type_codecs: NotRequired["list[dict[str, Any]]"] class _AsyncpgCloudSqlConnector: @@ -334,6 +340,9 @@ def __init__( self._user_connection_hook: Callable[[AsyncpgConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) + if connection_config and connection_config.get("pgbouncer"): + features_dict["pgbouncer"] = True + self._custom_type_codecs: list[dict[str, Any]] = list(features_dict.pop("type_codecs", None) or []) super().__init__( connection_config=build_connection_config(normalize_connection_config(connection_config)), @@ -457,6 +466,9 @@ async def _create_pool(self) -> "Pool[Record]": key: value for key, value in build_connection_config(self.connection_config).items() if value is not None } + if self.connection_config.get("pgbouncer") or self.driver_features.get("pgbouncer"): + config["statement_cache_size"] = 0 + if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): @@ -467,7 +479,7 @@ async def _create_pool(self) -> "Pool[Record]": return await asyncpg_create_pool(**config) async def _init_connection(self, connection: "AsyncpgConnection") -> None: - """Initialize connection with JSON codecs, pgvector support, and user callback. + """Initialize connection with JSON codecs, pgvector support, custom codecs, and user callback. Args: connection: AsyncPG connection to initialize. @@ -479,7 +491,6 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: decoder=self.driver_features.get("json_deserializer", from_json), ) - # Detect extensions on first connection, update dialect if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -500,7 +511,15 @@ async def _init_connection(self, connection: "AsyncpgConnection") -> None: if self._pgvector_available: await register_pgvector_support(connection) - # Call user-provided callback after internal setup + for codec in self._custom_type_codecs: + codec_kwargs: dict[str, Any] = { + "schema": codec.get("schema", "public"), + "format": codec.get("format", "text"), + } + codec_kwargs["encoder"] = codec["encoder"] + codec_kwargs["decoder"] = codec["decoder"] + await connection.set_type_codec(codec["typename"], **codec_kwargs) + if self._user_connection_hook is not None: await self._user_connection_hook(connection) @@ -544,6 +563,9 @@ async def create_connection(self) -> "AsyncpgConnection": for key in _POOL_ONLY_CONFIG_KEYS: config.pop(key, None) + if self.driver_features.get("pgbouncer"): + config["statement_cache_size"] = 0 + if self.driver_features.get("enable_cloud_sql", False): self._setup_cloud_sql_connector(config) elif self.driver_features.get("enable_alloydb", False): diff --git a/sqlspec/adapters/asyncpg/core.py b/sqlspec/adapters/asyncpg/core.py index ef7620600..94295400b 100644 --- a/sqlspec/adapters/asyncpg/core.py +++ b/sqlspec/adapters/asyncpg/core.py @@ -131,6 +131,9 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str if found_user: config["user"] = user_val + if config.pop("pgbouncer", False): + config["statement_cache_size"] = 0 + return config @@ -167,7 +170,7 @@ def configure_parameter_serializers( async def invoke_prepared_statement( - prepared: Any, parameters: "tuple[Any, ...] | dict[str, Any] | list[Any] | None", *, fetch: bool + prepared: Any, parameters: "tuple[Any, ...] | dict[str, Any] | list[Any] | None", *, fetch: bool, **kwargs: Any ) -> Any: """Invoke an AsyncPG prepared statement with optional parameters. @@ -175,25 +178,26 @@ async def invoke_prepared_statement( prepared: AsyncPG prepared statement object. parameters: Prepared parameters payload. fetch: Whether to fetch rows. + **kwargs: Native execution options, including timeout. Returns: Query result or status message. """ if parameters is None: if fetch: - return await prepared.fetch() - await prepared.fetch() + return await prepared.fetch(**kwargs) + await prepared.fetch(**kwargs) return prepared.get_statusmsg() if isinstance(parameters, dict): if fetch: - return await prepared.fetch(**parameters) - await prepared.fetch(**parameters) + return await prepared.fetch(**parameters, **kwargs) + await prepared.fetch(**parameters, **kwargs) return prepared.get_statusmsg() if fetch: - return await prepared.fetch(*parameters) - await prepared.fetch(*parameters) + return await prepared.fetch(*parameters, **kwargs) + await prepared.fetch(*parameters, **kwargs) return prepared.get_statusmsg() @@ -426,9 +430,17 @@ def collect_rows(records: "list[Any] | None") -> "tuple[list[Any], list[str]]": class AsyncpgStreamSource: """Compiled async chunk source streaming dict rows from an asyncpg cursor in a stream-owned transaction.""" - __slots__ = ("_chunk_size", "_cursor", "_driver", "_parameters", "_sql", "_transaction") - - def __init__(self, driver: Any, sql: str, parameters: "tuple[Any, ...]", chunk_size: int) -> None: + __slots__ = ("_chunk_size", "_cursor", "_driver", "_parameters", "_sql", "_timeout_args", "_transaction") + + def __init__( + self, + driver: Any, + sql: str, + parameters: "tuple[Any, ...]", + chunk_size: int, + timeout_args: "dict[str, Any] | None" = None, + ) -> None: + self._timeout_args = timeout_args or {} self._driver = driver self._sql = sql self._parameters = parameters @@ -446,7 +458,7 @@ async def _start(self) -> None: await transaction.start() self._transaction = transaction try: - self._cursor = await self._driver.connection.cursor(self._sql, *self._parameters) + self._cursor = await self._driver.connection.cursor(self._sql, *self._parameters, **self._timeout_args) except BaseException: await transaction.rollback() self._transaction = None @@ -454,7 +466,9 @@ async def _start(self) -> None: async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() - records = await self._driver._run_with_exception_handler(handler, self._cursor.fetch, self._chunk_size) + records = await self._driver._run_with_exception_handler( + handler, self._cursor.fetch, self._chunk_size, **self._timeout_args + ) self._driver._check_pending_exception(handler) assert records is not None return [dict(record) for record in records] diff --git a/sqlspec/adapters/asyncpg/driver.py b/sqlspec/adapters/asyncpg/driver.py index eccab9969..60d775352 100644 --- a/sqlspec/adapters/asyncpg/driver.py +++ b/sqlspec/adapters/asyncpg/driver.py @@ -114,6 +114,14 @@ def __init__( self._prepared_statements: OrderedDict[str, AsyncpgPreparedStatement] = OrderedDict() self._transaction: Any = None + def _timeout_arguments(self, config: "StatementConfig") -> "dict[str, Any]": + for options in (config.execution_args, self.statement_config.execution_args): + if options: + for key in ("timeout", "command_timeout"): + if key in options: + return {"timeout": options[key]} + return {} + async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") -> "ExecutionResult": """Execute single SQL statement. @@ -130,7 +138,7 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") params: tuple[Any, ...] = cast("tuple[Any, ...]", prepared_parameters) if prepared_parameters else () if statement.returns_rows(): - records = await cursor.fetch(sql, *params) if params else await cursor.fetch(sql) + records = await cursor.fetch(sql, *params, **self._timeout_arguments(statement.statement_config)) data, column_names = collect_rows(records) return self.create_execution_result( @@ -142,7 +150,7 @@ async def dispatch_execute(self, cursor: "AsyncpgConnection", statement: "SQL") row_format="record", ) - result = await cursor.execute(sql, *params) if params else await cursor.execute(sql) + result = await cursor.execute(sql, *params, **self._timeout_arguments(statement.statement_config)) affected_rows = parse_status(result) @@ -162,7 +170,7 @@ async def dispatch_execute_many(self, cursor: "AsyncpgConnection", statement: "S if prepared_parameters: parameter_sets = cast("list[Sequence[object]]", prepared_parameters) - await cursor.executemany(sql, parameter_sets) + await cursor.executemany(sql, parameter_sets, **self._timeout_arguments(statement.statement_config)) affected_rows = resolve_many_rowcount(parameter_sets) else: affected_rows = 0 @@ -186,7 +194,7 @@ async def dispatch_execute_script(self, cursor: "AsyncpgConnection", statement: last_result = None for stmt in statements: - result = await cursor.execute(stmt) + result = await cursor.execute(stmt, **self._timeout_arguments(statement.statement_config)) last_result = result successful_count += 1 @@ -301,7 +309,9 @@ def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "AsyncRow return None sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) params: tuple[Any, ...] = cast("tuple[Any, ...]", prepared_parameters) if prepared_parameters else () - return AsyncRowStream(AsyncpgStreamSource(self, sql, params, chunk_size)) + return AsyncRowStream( + AsyncpgStreamSource(self, sql, params, chunk_size, self._timeout_arguments(statement.statement_config)) + ) def handle_database_exceptions(self) -> "AsyncpgExceptionHandler": """Handle database exceptions with PostgreSQL error codes.""" @@ -310,7 +320,7 @@ def handle_database_exceptions(self) -> "AsyncpgExceptionHandler": async def execute_stack( self, stack: "StatementStack", *, continue_on_error: bool = False ) -> "tuple[StackResult, ...]": - """Execute a StatementStack using asyncpg's rapid batching.""" + """Execute a StatementStack sequentially, reusing native prepared statements.""" if not isinstance(stack, StatementStack) or not stack or self.stack_native_disabled: return await super().execute_stack(stack, continue_on_error=continue_on_error) @@ -509,7 +519,7 @@ async def _execute_stack_native( if not continue_on_error and not self._connection_in_transaction(): transaction_cm = self.connection.transaction() - with StackExecutionObserver(self, stack, continue_on_error, native_pipeline=True) as observer: + with StackExecutionObserver(self, stack, continue_on_error, native_pipeline=False) as observer: if transaction_cm is not None: async with transaction_cm: await self._run_stack_operations(stack, continue_on_error, observer, results) @@ -572,18 +582,22 @@ async def _run_stack_operations( await self._commit_stack_success() async def _execute_stack_operation_prepared(self, normalized: "NormalizedStackOperation") -> StackResult: - prepared = await self._get_prepared_statement(normalized.sql) + if self.driver_features.get("pgbouncer"): + result = await self._execute_cached_statement(normalized.statement) + return StackResult.from_sql_result(result) + timeout_args = self._timeout_arguments(normalized.statement.statement_config) + prepared = await self._get_prepared_statement(normalized.sql, timeout_args) metadata = {"prepared_statement": True} if normalized.statement.returns_rows(): - rows = await invoke_prepared_statement(prepared, normalized.parameters, fetch=True) + rows = await invoke_prepared_statement(prepared, normalized.parameters, fetch=True, **timeout_args) data, _ = collect_rows(rows) sql_result = create_sql_result( normalized.statement, data=data, rows_affected=len(data), metadata=metadata, row_format="record" ) return StackResult.from_sql_result(sql_result) - status = await invoke_prepared_statement(prepared, normalized.parameters, fetch=False) + status = await invoke_prepared_statement(prepared, normalized.parameters, fetch=False, **timeout_args) rowcount = parse_status(status) sql_result = create_sql_result(normalized.statement, rows_affected=rowcount, metadata=metadata) return StackResult.from_sql_result(sql_result) @@ -592,13 +606,13 @@ def _connection_in_transaction(self) -> bool: """Check if connection is in transaction.""" return bool(self.connection.is_in_transaction()) - async def _get_prepared_statement(self, sql: str) -> "AsyncpgPreparedStatement": + async def _get_prepared_statement(self, sql: str, timeout_args: "dict[str, Any]") -> "AsyncpgPreparedStatement": cached = self._prepared_statements.get(sql) if cached is not None: self._prepared_statements.move_to_end(sql) return cached - prepared = cast("AsyncpgPreparedStatement", await self.connection.prepare(sql)) + prepared = cast("AsyncpgPreparedStatement", await self.connection.prepare(sql, **timeout_args)) self._prepared_statements[sql] = prepared if len(self._prepared_statements) > PREPARED_STATEMENT_CACHE_SIZE: self._prepared_statements.popitem(last=False) diff --git a/sqlspec/adapters/cockroach_asyncpg/config.py b/sqlspec/adapters/cockroach_asyncpg/config.py index 158eb7942..ba762ea4d 100644 --- a/sqlspec/adapters/cockroach_asyncpg/config.py +++ b/sqlspec/adapters/cockroach_asyncpg/config.py @@ -7,7 +7,6 @@ from sqlspec.adapters.asyncpg.core import ( apply_driver_features, - build_connection_config, default_statement_config, register_json_codecs, register_pgvector_support, @@ -20,7 +19,7 @@ from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgRecord as Record from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_connect as asyncpg_connect from sqlspec.adapters.cockroach_asyncpg._typing import cockroach_asyncpg_create_pool as asyncpg_create_pool -from sqlspec.adapters.cockroach_asyncpg.core import validate_follower_read_staleness +from sqlspec.adapters.cockroach_asyncpg.core import build_connection_config, validate_follower_read_staleness from sqlspec.adapters.cockroach_asyncpg.driver import CockroachAsyncpgDriver, CockroachAsyncpgExceptionHandler from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs from sqlspec.core.capabilities import TypeCoercionCapabilities @@ -80,6 +79,9 @@ class CockroachAsyncpgConnectionConfig(TypedDict): timeout: NotRequired[float] connect_timeout: NotRequired[float] command_timeout: NotRequired[float] + application_name: NotRequired[str] + default_transaction_use_follower_reads: NotRequired[bool] + results_buffer_size: NotRequired[int] statement_cache_size: NotRequired[int] max_cached_statement_lifetime: NotRequired[int] max_cacheable_statement_size: NotRequired[int] diff --git a/sqlspec/adapters/cockroach_asyncpg/core.py b/sqlspec/adapters/cockroach_asyncpg/core.py index fefa80e03..0bd5210b6 100644 --- a/sqlspec/adapters/cockroach_asyncpg/core.py +++ b/sqlspec/adapters/cockroach_asyncpg/core.py @@ -8,6 +8,7 @@ from sqlglot import tokenize from sqlglot.tokenizer_core import TokenType +from sqlspec.adapters.asyncpg.core import build_connection_config as asyncpg_build_connection_config from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate @@ -19,6 +20,7 @@ __all__ = ( "CockroachAsyncpgRetryConfig", + "build_connection_config", "build_native_export", "build_native_import", "calculate_backoff_seconds", @@ -65,6 +67,31 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachAsyncp ) +def build_connection_config(config: "dict[str, Any]") -> "dict[str, Any]": + """Prepare CockroachDB AsyncPG connection config, extracting multi-region server settings.""" + result = asyncpg_build_connection_config(config) + server_settings = dict(result.get("server_settings") or {}) + if "application_name" in result: + server_settings.setdefault("application_name", str(result.pop("application_name"))) + if "default_transaction_use_follower_reads" in result: + val = result.pop("default_transaction_use_follower_reads") + if not isinstance(val, bool): + msg = "default_transaction_use_follower_reads must be a boolean" + raise ImproperConfigurationError(msg) + server_settings.setdefault("default_transaction_use_follower_reads", "on" if val else "off") + for key in ("results_buffer_size",): + if key not in result: + continue + value = result.pop(key) + if not isinstance(value, int) or isinstance(value, bool) or value < 0: + msg = f"{key} must be a non-negative integer" + raise ImproperConfigurationError(msg) + server_settings.setdefault(key, str(value)) + if server_settings: + result["server_settings"] = server_settings + return result + + def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. diff --git a/sqlspec/adapters/cockroach_psycopg/config.py b/sqlspec/adapters/cockroach_psycopg/config.py index 58b9c1ac5..a3b260d07 100644 --- a/sqlspec/adapters/cockroach_psycopg/config.py +++ b/sqlspec/adapters/cockroach_psycopg/config.py @@ -17,6 +17,7 @@ from sqlspec.adapters.cockroach_psycopg._typing import cockroach_psycopg_crdb as psycopg_crdb from sqlspec.adapters.cockroach_psycopg.core import ( apply_driver_features, + build_connection_config, build_statement_config, validate_follower_read_staleness, ) @@ -37,10 +38,9 @@ ) from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.events import EventRuntimeHints -from sqlspec.utils.config_tools import normalize_connection_config if TYPE_CHECKING: - from collections.abc import Awaitable, Callable, Mapping + from collections.abc import Awaitable, Callable from types import TracebackType from sqlspec.core import StatementConfig @@ -69,6 +69,10 @@ class CockroachPsycopgConnectionConfig(TypedDict): dbname: NotRequired[str] connect_timeout: NotRequired[int] options: NotRequired[str] + default_transaction_use_follower_reads: NotRequired[bool] + results_buffer_size: NotRequired[int] + statement_timeout: NotRequired[int] + idle_in_transaction_session_timeout: NotRequired[int] application_name: NotRequired[str] sslmode: NotRequired[str] sslcert: NotRequired[str] @@ -142,33 +146,6 @@ class CockroachPsycopgDriverFeatures(TypedDict): events_backend: NotRequired[Literal["poll_queue"]] -def build_connection_config( - connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", -) -> dict[str, Any]: - """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility. - - Maps connection string aliases (dsn, url, connection_string) to conninfo, database aliases - (database, db) to dbname, and user aliases (username) to user, while discarding redundant keys - that libpq rejects. - """ - config = normalize_connection_config(connection_config) - conninfo = ( - config.pop("conninfo", None) - or config.pop("dsn", None) - or config.pop("url", None) - or config.pop("connection_string", None) - ) - if conninfo is not None: - config["conninfo"] = conninfo - dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) - if dbname is not None: - config["dbname"] = dbname - user = config.pop("user", None) or config.pop("username", None) - if user is not None: - config["user"] = user - return config - - class CockroachPsycopgSyncConnectionContext(SyncPoolConnectionContext): """Context manager for CockroachDB psycopg connections.""" diff --git a/sqlspec/adapters/cockroach_psycopg/core.py b/sqlspec/adapters/cockroach_psycopg/core.py index 19dbc8257..7157e9726 100644 --- a/sqlspec/adapters/cockroach_psycopg/core.py +++ b/sqlspec/adapters/cockroach_psycopg/core.py @@ -10,17 +10,20 @@ from sqlspec.adapters.psycopg.core import apply_driver_features, build_statement_config, driver_profile from sqlspec.exceptions import ImproperConfigurationError, SerializationConflictError, SQLSpecError +from sqlspec.utils.config_tools import normalize_connection_config from sqlspec.utils.text import quote_identifier, split_qualified_identifier from sqlspec.utils.type_guards import has_sqlstate if TYPE_CHECKING: from collections.abc import Mapping + from sqlspec.adapters.cockroach_psycopg.config import CockroachPsycopgPoolConfig from sqlspec.storage import StorageTelemetry __all__ = ( "CockroachPsycopgRetryConfig", "apply_driver_features", + "build_connection_config", "build_native_export", "build_native_import", "build_statement_config", @@ -70,6 +73,49 @@ def from_features(cls, driver_features: "Mapping[str, Any]") -> "CockroachPsycop ) +def build_connection_config( + connection_config: "CockroachPsycopgPoolConfig | Mapping[str, Any] | None", +) -> dict[str, Any]: + """Build normalized CockroachDB psycopg connection configuration, resolving aliases for libpq compatibility.""" + config = normalize_connection_config(connection_config) + conninfo = ( + config.pop("conninfo", None) + or config.pop("dsn", None) + or config.pop("url", None) + or config.pop("connection_string", None) + ) + if conninfo is not None: + config["conninfo"] = conninfo + dbname = config.pop("dbname", None) or config.pop("database", None) or config.pop("db", None) + if dbname is not None: + config["dbname"] = dbname + user = config.pop("user", None) or config.pop("username", None) + if user is not None: + config["user"] = user + + session_options: list[str] = [] + if "default_transaction_use_follower_reads" in config: + val = config.pop("default_transaction_use_follower_reads") + if not isinstance(val, bool): + msg = "default_transaction_use_follower_reads must be a boolean" + raise ImproperConfigurationError(msg) + session_options.append(f"-c default_transaction_use_follower_reads={'on' if val else 'off'}") + for key in ("results_buffer_size", "statement_timeout", "idle_in_transaction_session_timeout"): + if key not in config: + continue + value = config.pop(key) + if not isinstance(value, int) or isinstance(value, bool) or value < 0: + msg = f"{key} must be a non-negative integer" + raise ImproperConfigurationError(msg) + session_options.append(f"-c {key}={value}") + if session_options: + existing = config.get("options") + opt_str = " ".join(session_options) + config["options"] = f"{existing} {opt_str}" if existing else opt_str + + return config + + def is_retryable_error(error: BaseException) -> bool: """Return True when the error should trigger a CockroachDB retry. diff --git a/sqlspec/adapters/psqlpy/core.py b/sqlspec/adapters/psqlpy/core.py index f525a7c50..9f764ef47 100644 --- a/sqlspec/adapters/psqlpy/core.py +++ b/sqlspec/adapters/psqlpy/core.py @@ -13,6 +13,7 @@ from sqlglot.errors import ParseError from sqlspec.adapters.psqlpy._typing import PsqlpyDataError, PsqlpyIntegrityError, PsqlpyOperationalError +from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector from sqlspec.core import ( DriverParameterProfile, ParameterStyle, @@ -596,6 +597,8 @@ def _coerce_parameter_for_cast(value: Any, cast_type: str, serializer: "Callable return _coerce_json_parameter(value, upper_cast, serializer) if upper_cast in _UUID_CASTS: return _coerce_uuid_parameter(value) + if upper_cast == "VECTOR": + return coerce_pgvector(value) if upper_cast in _TIMESTAMP_CASTS: return _coerce_timestamp_parameter(value) return value diff --git a/sqlspec/adapters/psqlpy/type_converter.py b/sqlspec/adapters/psqlpy/type_converter.py index 5f788c0d9..7d3941837 100644 --- a/sqlspec/adapters/psqlpy/type_converter.py +++ b/sqlspec/adapters/psqlpy/type_converter.py @@ -4,14 +4,28 @@ driver configuration layer. """ -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any -from sqlspec.typing import PGVECTOR_INSTALLED +from sqlspec.typing import PGVECTOR_INSTALLED, import_optional_attr if TYPE_CHECKING: from sqlspec.adapters.psqlpy._typing import PsqlpyConnection as Connection -__all__ = ("register_pgvector",) +__all__ = ("coerce_pgvector", "register_pgvector") + +_PGVECTOR_TYPE = import_optional_attr("psqlpy.extra_types", "PgVector") + + +def coerce_pgvector(value: Any) -> Any: + """Wrap supported dense vector values using psqlpy's native vector encoder.""" + if value is None or _PGVECTOR_TYPE is None or isinstance(value, _PGVECTOR_TYPE): + return value + if isinstance(value, (list, tuple)): + return _PGVECTOR_TYPE(list(value)) + tolist = getattr(value, "tolist", None) + if callable(tolist): + return _PGVECTOR_TYPE(tolist()) + return value def register_pgvector(connection: "Connection") -> None: diff --git a/sqlspec/adapters/psycopg/_typing.py b/sqlspec/adapters/psycopg/_typing.py index f59eb308c..b30aa323a 100644 --- a/sqlspec/adapters/psycopg/_typing.py +++ b/sqlspec/adapters/psycopg/_typing.py @@ -26,7 +26,9 @@ from psycopg.sql import Identifier as PsycopgIdentifier from psycopg.types.json import Jsonb as PsycopgJsonb from psycopg_pool import AsyncConnectionPool as PsycopgAsyncConnectionPool +from psycopg_pool import AsyncNullConnectionPool as PsycopgAsyncNullConnectionPool from psycopg_pool import ConnectionPool as PsycopgConnectionPool +from psycopg_pool import NullConnectionPool as PsycopgNullConnectionPool from psycopg_pool.abc import AsyncConnectFailedCB as PsycopgAsyncConnectFailedCB from psycopg_pool.abc import AsyncConnectionCB as PsycopgAsyncConnectionCB from psycopg_pool.abc import ConnectFailedCB as PsycopgConnectFailedCB @@ -65,6 +67,7 @@ "PsycopgAsyncConnectionCB", "PsycopgAsyncConnectionPool", "PsycopgAsyncCursor", + "PsycopgAsyncNullConnectionPool", "PsycopgAsyncRawCursor", "PsycopgAsyncRowFactory", "PsycopgAsyncSessionContext", @@ -79,6 +82,7 @@ "PsycopgJsonb", "PsycopgNativeAsyncConnection", "PsycopgNativeAsyncCursor", + "PsycopgNullConnectionPool", "PsycopgPipelineDriver", "PsycopgProgrammingError", "PsycopgRowFactory", diff --git a/sqlspec/adapters/psycopg/config.py b/sqlspec/adapters/psycopg/config.py index c8395c7e9..660314a61 100644 --- a/sqlspec/adapters/psycopg/config.py +++ b/sqlspec/adapters/psycopg/config.py @@ -4,6 +4,8 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypedDict, cast from mypy_extensions import mypyc_attr +from psycopg.adapt import AdaptersMap +from psycopg.types.json import set_json_dumps, set_json_loads from typing_extensions import NotRequired, Self from sqlspec.adapters.psycopg._typing import ( @@ -16,7 +18,9 @@ PsycopgSyncSessionContext, ) from sqlspec.adapters.psycopg._typing import PsycopgAsyncConnectionPool as AsyncConnectionPool +from sqlspec.adapters.psycopg._typing import PsycopgAsyncNullConnectionPool as AsyncNullConnectionPool from sqlspec.adapters.psycopg._typing import PsycopgConnectionPool as ConnectionPool +from sqlspec.adapters.psycopg._typing import PsycopgNullConnectionPool as NullConnectionPool from sqlspec.adapters.psycopg.core import apply_driver_features, default_statement_config from sqlspec.adapters.psycopg.driver import ( PsycopgAsyncDriver, @@ -43,6 +47,7 @@ from sqlspec.extensions.events import EventRuntimeHints from sqlspec.typing import ALLOYDB_CONNECTOR_INSTALLED from sqlspec.utils.config_tools import normalize_connection_config +from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Awaitable, Callable, Mapping @@ -86,6 +91,7 @@ "min_size", "name", "num_workers", + "null_pool", "open", "reconnect_failed", "reconnect_timeout", @@ -139,6 +145,7 @@ class PsycopgPoolParams(PsycopgConnectionParams): close_returns: NotRequired[bool] reconnect_failed: NotRequired["ConnectFailedCB | AsyncConnectFailedCB | None"] kwargs: NotRequired["dict[str, Any]"] + null_pool: NotRequired[bool] class PsycopgDriverFeatures(TypedDict): @@ -183,6 +190,7 @@ class PsycopgDriverFeatures(TypedDict): enable_alloydb_iam_auth: Enable AlloyDB IAM database authentication for sync connector connections. Defaults to False. alloydb_ip_type: AlloyDB connector IP type. Defaults to PRIVATE. + null_pool: Enable NullConnectionPool / AsyncNullConnectionPool for serverless / PgBouncer environments. """ enable_pgvector: NotRequired[bool] @@ -197,6 +205,7 @@ class PsycopgDriverFeatures(TypedDict): alloydb_instance_uri: NotRequired[str] enable_alloydb_iam_auth: NotRequired[bool] alloydb_ip_type: NotRequired[str] + null_pool: NotRequired[bool] def build_connection_config(connection_config: "PsycopgPoolParams | Mapping[str, Any] | None") -> dict[str, Any]: @@ -389,6 +398,7 @@ def _setup_alloydb_connector( def _create_pool(self) -> "ConnectionPool": """Create the actual connection pool.""" all_config = dict(self.connection_config) + all_config.pop("null_pool", None) pool_parameters = { "connection_class": all_config.pop("connection_class", None), @@ -422,10 +432,15 @@ def _create_pool(self) -> "ConnectionPool": self._setup_alloydb_connector(all_config, pool_parameters) conninfo = None + is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) + pool_cls = NullConnectionPool if is_null_pool else ConnectionPool + if is_null_pool: + pool_parameters.pop("min_size", None) + if conninfo: - pool = ConnectionPool(conninfo, kwargs=all_config, **pool_parameters) + pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) else: - pool = ConnectionPool("", kwargs=all_config, **pool_parameters) + pool = pool_cls("", kwargs=all_config, **pool_parameters) return pool @@ -434,7 +449,12 @@ def _configure_connection(self, conn: "PsycopgSyncConnection") -> None: if autocommit_setting is not None: conn.autocommit = autocommit_setting - # Detect extensions on first connection, update dialect + serializer = self.driver_features.get("json_serializer", to_json) + deserializer = self.driver_features.get("json_deserializer", from_json) + if isinstance(getattr(conn, "adapters", None), AdaptersMap): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) + if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) @@ -711,6 +731,7 @@ async def _create_pool(self) -> "AsyncConnectionPool": """Create the actual async connection pool.""" all_config = dict(self.connection_config) + all_config.pop("null_pool", None) pool_parameters = { "connection_class": all_config.pop("connection_class", None), @@ -738,10 +759,16 @@ async def _create_pool(self) -> "AsyncConnectionPool": conninfo = all_config.pop("conninfo", None) kwargs = all_config.pop("kwargs", {}) all_config.update(kwargs) + + is_null_pool = bool(self.connection_config.get("null_pool") or self.driver_features.get("null_pool")) + pool_cls = AsyncNullConnectionPool if is_null_pool else AsyncConnectionPool + if is_null_pool: + pool_parameters.pop("min_size", None) + if conninfo: - pool = AsyncConnectionPool(conninfo, kwargs=all_config, **pool_parameters) + pool = pool_cls(conninfo, kwargs=all_config, **pool_parameters) else: - pool = AsyncConnectionPool("", kwargs=all_config, **pool_parameters) + pool = pool_cls("", kwargs=all_config, **pool_parameters) if open_pool is True: await pool.open() @@ -753,7 +780,12 @@ async def _configure_async_connection(self, conn: "PsycopgAsyncConnection") -> N if autocommit_setting is not None: await conn.set_autocommit(autocommit_setting) - # Detect extensions on first connection, update dialect + serializer = self.driver_features.get("json_serializer", to_json) + deserializer = self.driver_features.get("json_deserializer", from_json) + if isinstance(getattr(conn, "adapters", None), AdaptersMap): + set_json_dumps(serializer, conn) + set_json_loads(deserializer, conn) + if self._pgvector_available is None: detected_extensions: set[str] = set() extensions = build_postgres_extension_probe_names(self.driver_features) diff --git a/tests/unit/adapters/test_postgres_native_options.py b/tests/unit/adapters/test_postgres_native_options.py new file mode 100644 index 000000000..8f76f8d18 --- /dev/null +++ b/tests/unit/adapters/test_postgres_native_options.py @@ -0,0 +1,127 @@ +"""Native PostgreSQL option forwarding without database services.""" + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from asyncpg import Connection + +from sqlspec.adapters.asyncpg.config import AsyncpgConfig +from sqlspec.adapters.asyncpg.driver import AsyncpgDriver +from sqlspec.adapters.cockroach_asyncpg.core import build_connection_config as asyncpg_cockroach_config +from sqlspec.adapters.cockroach_psycopg.core import build_connection_config as psycopg_cockroach_config +from sqlspec.adapters.psycopg import config as psycopg_module +from sqlspec.adapters.psycopg.config import PsycopgAsyncConfig, PsycopgSyncConfig +from sqlspec.core import StatementConfig, StatementStack +from sqlspec.exceptions import ImproperConfigurationError + + +async def test_asyncpg_timeout_survives_repeated_execution() -> None: + connection = MagicMock(spec=Connection) + connection.fetch = AsyncMock(return_value=[]) + driver = AsyncpgDriver( + connection, statement_config=StatementConfig(dialect="postgres", execution_args={"timeout": 0}) + ) + for _ in range(2): + await driver.execute("SELECT 1") + assert [call.kwargs for call in connection.fetch.call_args_list] == [{"timeout": 0}, {"timeout": 0}] + + +async def test_asyncpg_custom_codec_is_registered_after_internal_setup() -> None: + encoder = str + decoder = int + config = AsyncpgConfig( + driver_features={"type_codecs": [{"typename": "custom", "encoder": encoder, "decoder": decoder}]} + ) + config._pgvector_available = False + connection = MagicMock(spec=Connection) + connection.set_type_codec = AsyncMock() + await config._init_connection(connection) + assert connection.set_type_codec.call_args.args == ("custom",) + assert connection.set_type_codec.call_args.kwargs == { + "schema": "public", + "format": "text", + "encoder": encoder, + "decoder": decoder, + } + + +def test_asyncpg_pgbouncer_disables_prepared_cache_consistently() -> None: + config = AsyncpgConfig(connection_config={"pgbouncer": True, "statement_cache_size": 128}) + assert config.connection_config["statement_cache_size"] == 0 + assert config.driver_features["pgbouncer"] is True + + +@pytest.mark.parametrize( + "config_type,pool_name", + [(PsycopgSyncConfig, "NullConnectionPool"), (PsycopgAsyncConfig, "AsyncNullConnectionPool")], +) +async def test_null_pool_preserves_throttling_without_leaking_flag( + monkeypatch: pytest.MonkeyPatch, config_type: Any, pool_name: str +) -> None: + constructor = MagicMock() + monkeypatch.setattr(psycopg_module, pool_name, constructor) + config = config_type(connection_config={"null_pool": True, "open": False, "max_size": 7, "max_waiting": 2}) + if config_type is PsycopgAsyncConfig: + await config._create_pool() + else: + config._create_pool() + args = constructor.call_args.kwargs + assert args["max_size"] == 7 + assert args["max_waiting"] == 2 + assert "null_pool" not in args["kwargs"] + assert "null_pool" not in config._connection_kwargs()[2] + + +@pytest.mark.parametrize("builder", [asyncpg_cockroach_config, psycopg_cockroach_config]) +def test_cockroach_options_reject_invalid_values(builder: Any) -> None: + with pytest.raises(ImproperConfigurationError, match="boolean"): + builder({"default_transaction_use_follower_reads": "invalid"}) + with pytest.raises(ImproperConfigurationError, match="integer"): + builder({"results_buffer_size": "1024 -c injected=1"}) + + +def test_cockroach_options_preserve_native_controls() -> None: + async_config = asyncpg_cockroach_config({ + "application_name": "test", + "default_transaction_use_follower_reads": True, + "results_buffer_size": 1024, + }) + assert async_config["server_settings"] == { + "application_name": "test", + "default_transaction_use_follower_reads": "on", + "results_buffer_size": "1024", + } + sync_config = psycopg_cockroach_config({ + "options": "-c application_name=test", + "statement_timeout": 100, + "results_buffer_size": 1024, + }) + assert sync_config["options"] == "-c application_name=test -c results_buffer_size=1024 -c statement_timeout=100" + + +def test_psqlpy_dense_vector_uses_native_wrapper() -> None: + from psqlpy.extra_types import PgVector + + from sqlspec.adapters.psqlpy.type_converter import coerce_pgvector + + vector = coerce_pgvector([1.0, 2.0]) + assert isinstance(vector, PgVector) + assert coerce_pgvector(vector) is vector + + +async def test_pgbouncer_stack_cancellation_keeps_transaction_cleanup() -> None: + connection = MagicMock(spec=Connection) + connection.is_in_transaction.return_value = False + connection.fetch = AsyncMock(side_effect=asyncio.CancelledError) + transaction = MagicMock() + transaction.__aenter__ = AsyncMock() + transaction.__aexit__ = AsyncMock(return_value=False) + connection.transaction.return_value = transaction + driver = AsyncpgDriver(connection, driver_features={"pgbouncer": True}) + stack = StatementStack().push_execute("SELECT 1") + with pytest.raises(asyncio.CancelledError): + await driver.execute_stack(stack) + assert transaction.__aexit__.call_args.args[0] is asyncio.CancelledError + connection.prepare.assert_not_called() From 7d0e2b622db5fb53e446046bb02ea56b51b4d459 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:05:47 +0000 Subject: [PATCH 09/10] docs: reconcile unreleased adapter changelog after rebase --- docs/changelog.rst | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index a5281a602..d8555cc36 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,13 +1,13 @@ -== +========= Changelog -== +========= All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates -======= +============== Unreleased ---------- @@ -45,6 +45,12 @@ Unreleased **Fixed:** +* 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. + +* Asyncpg stack telemetry reports sequential prepared execution rather than + native pipelining. Each statement still returns its own result. + * Builder upserts emit ``MERGE`` for the ``db2`` dialect. * ADBC ADK stores reuse cached PostgreSQL placeholder conversion and preserve @@ -2237,7 +2243,7 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions -========== +================= For releases before ``v0.24.0``, see the repository tag history and GitHub release records. From 7463ad156d78d6ed94d80b2d010e626d8d708588 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 28 Sep 2026 00:33:15 +0000 Subject: [PATCH 10/10] test: register psycopg null pool feature consumption --- 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 0f3480561..66bb8a7ca 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -219,6 +219,7 @@ class SourceEquivalenceCase: "alloydb_instance_uri", "enable_alloydb_iam_auth", "alloydb_ip_type", + "null_pool", ), "pymssql": ("json_serializer", "json_deserializer", "on_connection_create", "enable_events", "events_backend"), "pymysql": (