Source code for sqlspec.adapters.cockroach_asyncpg.driver

"""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
async def run_transaction_with_retry(self, operation: "Callable[[], Awaitable[_T]]") -> _T: """Execute a full CockroachDB transaction callback with serialization retries.""" if not self._enable_retry or self._connection_in_transaction(): return await operation() last_error: BaseException | None = None for attempt in range(self._retry_config.max_retries + 1): try: await self.begin() result = await operation() await self.commit() except Exception as exc: last_error = exc with contextlib.suppress(Exception): await self.rollback() if not is_retryable_error(exc) or attempt >= self._retry_config.max_retries: raise else: return result delay = calculate_backoff_seconds(attempt, self._retry_config) if self._retry_config.enable_logging: logger.debug("CockroachDB retry %s/%s after %.3fs", attempt + 1, self._retry_config.max_retries, delay) await asyncio.sleep(delay) msg = "CockroachDB transaction retry limit exceeded" raise TransactionRetryError(msg) from last_error