diff --git a/docs/changelog.rst b/docs/changelog.rst index 3174eff84..d8555cc36 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -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,14 @@ 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 question marks in quoted identifiers, literals, and comments. @@ -45,7 +60,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. 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/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..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/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..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/adk/store.py b/sqlspec/adapters/cockroach_psycopg/adk/store.py index fc868818e..8d81ec748 100644 --- a/sqlspec/adapters/cockroach_psycopg/adk/store.py +++ b/sqlspec/adapters/cockroach_psycopg/adk/store.py @@ -229,24 +229,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) + 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 @@ -705,24 +715,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) + 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 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/adk/store.py b/sqlspec/adapters/psqlpy/adk/store.py index ecbe6df2f..d9503c6e9 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,36 @@ 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]) + 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 """ - await conn.execute(sql, [session_id, app_name, user_id, state]) + result = await conn.fetch(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." + rows: list[dict[str, Any]] = result.result() if result else [] + if not rows: + msg = "Failed to fetch created session" raise RuntimeError(msg) - return res + row = rows[0] + 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,7 +142,7 @@ async def get_session( """ 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, session_id]) rows: list[dict[str, Any]] = result.result() if result else [] @@ -158,7 +170,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 +208,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 +231,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 +241,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 +292,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 +362,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 +394,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 +416,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 +440,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 +455,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 +468,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 +486,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 +498,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 +521,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 +632,8 @@ class PsqlpyADKMemoryStore(BaseAsyncADKMemoryStore["PsqlpyConfig"]): __slots__ = () + _config: "PsqlpyConfig" + def __init__(self, config: "PsqlpyConfig") -> None: """Initialize Psqlpy memory store.""" super().__init__(config) @@ -668,7 +682,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 +741,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 +757,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 +784,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 +863,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 +897,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..dd99bc8e1 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__( @@ -285,11 +297,10 @@ 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: - return - if self._user_connection_hook is not None: - await self._user_connection_hook(connection) - self._initialized_connection_ids.add(conn_id) + 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) async def _create_pool(self) -> "ConnectionPool": """Create the actual async connection pool.""" @@ -304,6 +315,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/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/litestar/store.py b/sqlspec/adapters/psqlpy/litestar/store.py index e4bbaf0f0..ff40e16f7 100644 --- a/sqlspec/adapters/psqlpy/litestar/store.py +++ b/sqlspec/adapters/psqlpy/litestar/store.py @@ -82,32 +82,34 @@ 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. """ + 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, + 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: + query_result = await conn.fetch(sql, [new_expires_at, key]) + rows = query_result.result() if query_result else [] + if not rows: + return None + return bytes(rows[0]["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() - + rows = query_result.result() if query_result else [] if not rows: 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"]) + 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/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/sqlspec/adapters/psycopg/driver.py b/sqlspec/adapters/psycopg/driver.py index 8c6b36f09..7f48b38f2 100644 --- a/sqlspec/adapters/psycopg/driver.py +++ b/sqlspec/adapters/psycopg/driver.py @@ -336,18 +336,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,21 +516,19 @@ 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: + 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)) - for record in prepared_records: + for record in records: copy_ctx.write_row(record) if exc_handler.pending_exception is not None: raise exc_handler.pending_exception from None @@ -853,18 +851,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,21 +1036,19 @@ 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: + 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)) - for record in prepared_records: + 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 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_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() 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."""