"""CockroachDB AsyncPG driver implementation."""
import asyncio
import contextlib
from typing import TYPE_CHECKING, Any, TypeVar, cast
from sqlspec.adapters.asyncpg.core import create_mapped_exception, driver_profile
from sqlspec.adapters.asyncpg.driver import AsyncpgDriver
from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgPostgresError, CockroachAsyncpgSessionContext
from sqlspec.adapters.cockroach_asyncpg.core import (
CockroachAsyncpgRetryConfig,
calculate_backoff_seconds,
is_retryable_error,
)
from sqlspec.adapters.cockroach_asyncpg.data_dictionary import CockroachAsyncpgDataDictionary
from sqlspec.core import SQL, register_driver_profile
from sqlspec.driver import BaseAsyncExceptionHandler
from sqlspec.exceptions import SerializationConflictError, TransactionRetryError
from sqlspec.utils.logging import get_logger
from sqlspec.utils.type_guards import has_sqlstate
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from sqlspec.adapters.cockroach_asyncpg._typing import CockroachAsyncpgConnection
from sqlspec.core import StatementConfig
from sqlspec.driver import ExecutionResult
__all__ = ("CockroachAsyncpgDriver", "CockroachAsyncpgExceptionHandler", "CockroachAsyncpgSessionContext")
logger = get_logger("sqlspec.adapters.cockroach_asyncpg")
_T = TypeVar("_T")
class CockroachAsyncpgExceptionHandler(BaseAsyncExceptionHandler):
"""Async context manager for CockroachDB AsyncPG exceptions."""
__slots__ = ()
def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool:
_ = exc_type
if isinstance(exc_val, CockroachAsyncpgPostgresError) or has_sqlstate(exc_val):
if has_sqlstate(exc_val) and str(exc_val.sqlstate) == "40001":
self.pending_exception = SerializationConflictError(str(exc_val))
return True
self.pending_exception = create_mapped_exception(exc_val)
return True
return False
class CockroachAsyncpgDriver(AsyncpgDriver):
"""CockroachDB AsyncPG driver with retry support."""
__slots__ = ("_enable_retry", "_follower_staleness", "_retry_config")
dialect = "postgres"
def __init__(
self,
connection: "CockroachAsyncpgConnection",
statement_config: "StatementConfig | None" = None,
driver_features: "dict[str, Any] | None" = None,
) -> None:
super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features)
self._retry_config = CockroachAsyncpgRetryConfig.from_features(self.driver_features)
self._enable_retry = bool(self.driver_features.get("enable_auto_retry", True))
self._follower_staleness = cast("str | None", self.driver_features.get("default_staleness"))
# Data dictionary is lazily initialized in property; use parent slot
self._data_dictionary = None