Source code for sqlspec.adapters.cockroach_asyncpg.config

"""CockroachDB AsyncPG configuration."""

from contextlib import suppress
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypedDict, cast

from typing_extensions import NotRequired

from sqlspec.adapters.asyncpg.core import (
    apply_driver_features,
    default_statement_config,
    register_json_codecs,
    register_pgvector_support,
)
from sqlspec.adapters.cockroach_asyncpg._typing import (
    CockroachAsyncpgConnection,
    CockroachAsyncpgPool,
    CockroachAsyncpgSessionContext,
)
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.driver import CockroachAsyncpgDriver, CockroachAsyncpgExceptionHandler
from sqlspec.config import AsyncDatabaseConfig, ExtensionConfigs
from sqlspec.core.capabilities import TypeCoercionCapabilities
from sqlspec.core.config_runtime import resolve_runtime_statement_config
from sqlspec.driver import AsyncPoolConnectionContext, AsyncPoolSessionFactory
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.extensions.events import EventRuntimeHints
from sqlspec.utils.config_tools import normalize_connection_config
from sqlspec.utils.serializers import from_json, to_json

if TYPE_CHECKING:
    from asyncio.events import AbstractEventLoop
    from collections.abc import Awaitable, Callable

    from sqlspec.core import StatementConfig
    from sqlspec.observability import ObservabilityConfig

__all__ = (
    "CockroachAsyncpgConfig",
    "CockroachAsyncpgConnectionConfig",
    "CockroachAsyncpgDriverFeatures",
    "CockroachAsyncpgGSSLib",
    "CockroachAsyncpgPoolConfig",
    "CockroachAsyncpgTargetSessionAttrs",
    "build_connection_config",
)

_POOL_ONLY_CONFIG_KEYS: Final[frozenset[str]] = frozenset({
    "init",
    "max_inactive_connection_lifetime",
    "max_queries",
    "max_size",
    "min_size",
    "reset",
    "setup",
})


CockroachAsyncpgTargetSessionAttrs = Literal["any", "primary", "standby", "read-write", "read-only", "prefer-standby"]
CockroachAsyncpgGSSLib = Literal["gssapi", "sspi"]


[docs] class CockroachAsyncpgConnectionConfig(TypedDict): """AsyncPG connection parameters for CockroachDB.""" dsn: NotRequired[str] host: NotRequired[str] port: NotRequired[int] user: NotRequired[str] password: NotRequired[str] service: NotRequired[str] servicefile: NotRequired[str] database: NotRequired[str] ssl: NotRequired[Any] passfile: NotRequired[str] direct_tls: NotRequired[bool] 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] server_settings: NotRequired["dict[str, str]"] target_session_attrs: NotRequired[CockroachAsyncpgTargetSessionAttrs] krbsrvname: NotRequired[str] gsslib: NotRequired[CockroachAsyncpgGSSLib]
[docs] class CockroachAsyncpgPoolConfig(CockroachAsyncpgConnectionConfig): """AsyncPG pool parameters for CockroachDB.""" min_size: NotRequired[int] max_size: NotRequired[int] max_queries: NotRequired[int] max_inactive_connection_lifetime: NotRequired[float] connect: NotRequired["Callable[..., Awaitable[CockroachAsyncpgConnection]]"] setup: NotRequired["Callable[[CockroachAsyncpgConnection], Awaitable[None]]"] init: NotRequired["Callable[[CockroachAsyncpgConnection], Awaitable[None]]"] reset: NotRequired["Callable[[CockroachAsyncpgConnection], Awaitable[None]]"] loop: NotRequired["AbstractEventLoop"] connection_class: NotRequired[type["CockroachAsyncpgConnection"]] record_class: NotRequired[type[Record]] extra: NotRequired["dict[str, Any]"]
class _NativeStorageCSVOptions(TypedDict): """Explicit CSV conventions; no header or NULL marker is inferred.""" nullas: NotRequired[str] nullif: NotRequired[str] skip: NotRequired[int] def _validate_follower_read_features(driver_features: "dict[str, Any]") -> None: staleness = driver_features.get("default_staleness") if staleness is None: return if not isinstance(staleness, str): msg = "default_staleness must be a string." raise ImproperConfigurationError(msg) driver_features["default_staleness"] = validate_follower_read_staleness(staleness) def _validate_native_storage_options(driver_features: "dict[str, Any]") -> None: if "native_storage_csv_options" not in driver_features: return options = driver_features["native_storage_csv_options"] if not isinstance(options, dict) or options.keys() - {"nullas", "nullif", "skip"}: msg = "native_storage_csv_options must contain only nullas, nullif, and skip." raise ImproperConfigurationError(msg) for key in ("nullas", "nullif"): if key in options and not isinstance(options[key], str): msg = "native_storage_csv_options nullas and nullif must be strings." raise ImproperConfigurationError(msg) if "skip" in options and (type(options["skip"]) is not int or options["skip"] < 0): msg = "native_storage_csv_options skip must be a nonnegative integer, excluding bool." raise ImproperConfigurationError(msg) driver_features["native_storage_csv_options"] = dict(options)
[docs] class CockroachAsyncpgDriverFeatures(TypedDict): """Driver feature flags for CockroachDB AsyncPG adapter. enable_native_storage: Opt in to server-side storage operations. Exports create generated files under a prefix; imports temporarily take the table offline. native_storage_csv_options: Explicit nullas/nullif markers and skip count. Export is headerless; native CSV import requires an explicit skip count. Markers must not occur as literal data. No NULL marker is inferred. on_connection_create: Async callback executed when a connection is acquired from pool. Receives the raw asyncpg connection for low-level driver configuration. Called after internal setup (JSON codecs, pgvector registration). """ enable_native_storage: NotRequired[bool] native_storage_csv_options: NotRequired[_NativeStorageCSVOptions] enable_auto_retry: NotRequired[bool] max_retries: NotRequired[int] retry_delay_base_ms: NotRequired[float] retry_delay_max_ms: NotRequired[float] enable_retry_logging: NotRequired[bool] enable_follower_reads: NotRequired[bool] default_staleness: NotRequired[str] json_serializer: NotRequired["Callable[[Any], str]"] json_deserializer: NotRequired["Callable[[str], Any]"] enable_json_codecs: NotRequired[bool] enable_pgvector: NotRequired[bool] on_connection_create: "NotRequired[Callable[[CockroachAsyncpgConnection], Awaitable[None]]]" enable_events: NotRequired[bool] events_backend: NotRequired[Literal["poll_queue"]]
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: super().__init__(config) self._ctx: Any | None = None async def acquire_connection(self) -> "CockroachAsyncpgConnection": pool = self._config.connection_instance if pool is None: pool = await self._config.create_pool() self._config.connection_instance = pool ctx = pool.acquire() self._ctx = ctx return cast("CockroachAsyncpgConnection", await ctx.__aenter__()) async def release_connection(self, _conn: "CockroachAsyncpgConnection", **kwargs: Any) -> None: if self._ctx is not None: await self._ctx.__aexit__(kwargs.get("exc_type"), kwargs.get("exc_val"), kwargs.get("exc_tb")) self._ctx = None class CockroachAsyncpgConnectionContext(AsyncPoolConnectionContext): """Async context manager for CockroachDB AsyncPG connections.""" __slots__ = ()
[docs] class CockroachAsyncpgConfig( AsyncDatabaseConfig[CockroachAsyncpgConnection, CockroachAsyncpgPool, CockroachAsyncpgDriver] ): """Configuration for CockroachDB using AsyncPG.""" driver_type: "ClassVar[type[CockroachAsyncpgDriver]]" = CockroachAsyncpgDriver 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 supports_native_arrow_import: "ClassVar[bool]" = True supports_native_parquet_export: "ClassVar[bool]" = True supports_native_parquet_import: "ClassVar[bool]" = True supports_native_row_streaming: "ClassVar[bool]" = True type_coercion_capabilities: "ClassVar[TypeCoercionCapabilities]" = TypeCoercionCapabilities( datetime_binding="native", timestamp_precision="microsecond", json_columns_decoded=True, uuid_binding="native" ) _connection_context_class: "ClassVar[type[CockroachAsyncpgConnectionContext]]" = CockroachAsyncpgConnectionContext _session_factory_class: "ClassVar[type[_CockroachAsyncpgSessionFactory]]" = _CockroachAsyncpgSessionFactory _session_context_class: "ClassVar[type[CockroachAsyncpgSessionContext]]" = CockroachAsyncpgSessionContext _default_statement_config = default_statement_config
[docs] def __init__( self, *, connection_config: "CockroachAsyncpgPoolConfig | dict[str, Any] | None" = None, connection_instance: "CockroachAsyncpgPool | None" = None, migration_config: "dict[str, Any] | None" = None, statement_config: "StatementConfig | None" = None, driver_features: "CockroachAsyncpgDriverFeatures | dict[str, Any] | None" = None, bind_key: "str | None" = None, extension_config: "ExtensionConfigs | None" = None, observability_config: "ObservabilityConfig | None" = None, **kwargs: Any, ) -> None: raw_enable_pgvector = bool(driver_features and driver_features.get("enable_pgvector") is True) connection_config = build_connection_config(normalize_connection_config(connection_config)) statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, driver_features) driver_features["enable_pgvector"] = raw_enable_pgvector driver_features.setdefault("enable_native_storage", False) _validate_native_storage_options(driver_features) _validate_follower_read_features(driver_features) driver_features.setdefault("enable_auto_retry", True) features_dict = dict(driver_features) self._user_connection_hook: Callable[[CockroachAsyncpgConnection], Awaitable[None]] | None = features_dict.pop( "on_connection_create", None ) super().__init__( connection_config=connection_config, connection_instance=connection_instance, migration_config=migration_config, statement_config=statement_config, driver_features=features_dict, bind_key=bind_key, extension_config=extension_config, observability_config=observability_config, **kwargs, )
async def _create_pool(self) -> "CockroachAsyncpgPool": config = build_connection_config(self.connection_config) config.setdefault("init", self._init_connection) return await asyncpg_create_pool(**config) async def _init_connection(self, connection: "CockroachAsyncpgConnection") -> None: """Initialize connection with JSON codecs and user callback.""" if self.driver_features.get("enable_json_codecs", True): await register_json_codecs( connection, encoder=self.driver_features.get("json_serializer", to_json), decoder=self.driver_features.get("json_deserializer", from_json), ) if self.driver_features.get("enable_pgvector") is True: await register_pgvector_support(connection) if self._user_connection_hook is not None: await self._user_connection_hook(connection) async def _close_pool(self) -> None: if not self.connection_instance: return await self.connection_instance.close() self.connection_instance = None
[docs] async def create_connection(self) -> "CockroachAsyncpgConnection": """Open a standalone connection owned by the caller. The connection carries the same connection settings and init hook the pool applies, consumes no pool slot, and must be closed by the caller. Returns: A CockroachDB asyncpg connection. """ config = build_connection_config(self.connection_config) for key in _POOL_ONLY_CONFIG_KEYS: config.pop(key, None) connect = config.pop("connect", None) connection = await connect() if connect is not None else await asyncpg_connect(**config) init = self.connection_config.get("init", self._init_connection) try: await init(connection) except BaseException: with suppress(Exception): await connection.close() raise return cast("CockroachAsyncpgConnection", connection)
[docs] def provide_session( self, *_args: Any, statement_config: "StatementConfig | None" = None, follower_reads: bool | None = None, staleness: str | None = None, **_kwargs: Any, ) -> "CockroachAsyncpgSessionContext": factory = _CockroachAsyncpgSessionFactory(self) driver_features = dict(self.driver_features) if follower_reads is not None: driver_features["enable_follower_reads"] = follower_reads if staleness is not None: driver_features["default_staleness"] = validate_follower_read_staleness(staleness) return CockroachAsyncpgSessionContext( acquire_connection=factory.acquire_connection, release_connection=factory.release_connection, statement_config=statement_config or (lambda: resolve_runtime_statement_config(None, self.statement_config, default_statement_config)), driver_features=driver_features, prepare_driver=self._prepare_driver, )
[docs] async def provide_pool(self, *args: Any, **kwargs: Any) -> "CockroachAsyncpgPool": if not self.connection_instance: self.connection_instance = await self.create_pool() return self.connection_instance
[docs] def get_signature_namespace(self) -> "dict[str, Any]": namespace = super().get_signature_namespace() namespace.update({ "CockroachAsyncpgConnectionConfig": CockroachAsyncpgConnectionConfig, "CockroachAsyncpgPoolConfig": CockroachAsyncpgPoolConfig, "CockroachAsyncpgDriver": CockroachAsyncpgDriver, "CockroachAsyncpgExceptionHandler": CockroachAsyncpgExceptionHandler, "CockroachAsyncpgSessionContext": CockroachAsyncpgSessionContext, }) return namespace
[docs] def get_event_runtime_hints(self) -> "EventRuntimeHints": return EventRuntimeHints(poll_interval=0.5, select_for_update=True, skip_locked=True)