"""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)