Source code for sqlspec.adapters.spanner.driver

"""Spanner driver implementation."""

import contextlib
from collections.abc import Iterator
from itertools import islice
from typing import TYPE_CHECKING, Any, Protocol, cast, overload

import sqlglot as _sqlglot
from sqlglot import exp as _sqlglot_exp

from sqlspec.adapters.spanner._typing import (
    SpannerGoogleAPICallError,
    SpannerSessionContext,
    SpannerSyncCursor,
    SpannerTransaction,
)
from sqlspec.adapters.spanner.core import (
    build_param_type_signature,
    coerce_params,
    collect_rows,
    create_mapped_exception,
    default_statement_config,
    driver_profile,
    infer_param_types,
    resolve_row_plan,
    supports_batch_update,
    supports_write,
)
from sqlspec.adapters.spanner.data_dictionary import SpannerDataDictionary
from sqlspec.core import StatementConfig, register_driver_profile
from sqlspec.driver import (
    BaseSyncExceptionHandler,
    ExecutionResult,
    SyncDriverAdapterBase,
    SyncRowStream,
    rows_to_dicts,
)
from sqlspec.exceptions import SQLConversionError
from sqlspec.utils.serializers import from_json

if TYPE_CHECKING:
    from collections.abc import Callable, Sequence

    from google.api_core.retry import Retry
    from google.cloud.spanner_v1 import DirectedReadOptions, RequestOptions
    from sqlglot.dialects.dialect import DialectType

    from sqlspec.adapters.spanner._typing import SpannerConnection
    from sqlspec.builder import QueryBuilder
    from sqlspec.core import ArrowResult, SQLResult, Statement, StatementFilter
    from sqlspec.core.statement import SQL
    from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry
    from sqlspec.typing import SchemaT, StatementParameters

__all__ = (
    "SpannerDataDictionary",
    "SpannerExceptionHandler",
    "SpannerSessionContext",
    "SpannerSyncCursor",
    "SpannerSyncDriver",
)

_READ_ONLY_SNAPSHOT_ERROR_MESSAGE = (
    "Cannot execute DML in a read-only Snapshot context. "
    "SpannerSyncConfig.provide_session() opens a write-capable Transaction by default; "
    "the current session must have been opened via SpannerSyncConfig.provide_read_session()."
)


class SpannerExceptionHandler(BaseSyncExceptionHandler):
    """Map Spanner client exceptions to SQLSpec exceptions.

    Uses deferred exception pattern for mypyc compatibility: exceptions
    are stored in pending_exception rather than raised from __exit__
    to avoid ABI boundary violations with compiled code.
    """

    __slots__ = ()

    def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "BaseException") -> bool:

        if exc_type is None:
            return False

        if isinstance(exc_val, SpannerGoogleAPICallError):
            self.pending_exception = create_mapped_exception(exc_val)
            return True
        return False


class SpannerSyncDriver(SyncDriverAdapterBase):
    """Synchronous Spanner driver operating on Snapshot or Transaction contexts."""

    dialect: "DialectType" = "spanner"
    __slots__ = ("_data_dictionary", "_pending_execute_options", "_row_plan_cache")

    def __init__(
        self,
        connection: "SpannerConnection",
        statement_config: "StatementConfig | None" = None,
        driver_features: "dict[str, Any] | None" = None,
    ) -> None:
        features = dict(driver_features) if driver_features else {}
        if statement_config is None:
            statement_config = default_statement_config

        super().__init__(connection=connection, statement_config=statement_config, driver_features=features)
        self._data_dictionary: SpannerDataDictionary | None = None
        self._pending_execute_options: _PerCallExecuteOptions | None = None
        self._row_plan_cache: dict[int, tuple[Any, list[str], tuple[tuple[int, Any], ...] | None]] = {}
# ───────────────────────────────────────────────────────────────────────────── # CORE DISPATCH METHODS - The Execution Engine # ───────────────────────────────────────────────────────────────────────────── def dispatch_execute(self, cursor: "SpannerConnection", statement: "SQL") -> ExecutionResult: sql, params = self._compiled_sql(statement, self.statement_config) params = cast("dict[str, Any] | None", params) coerced_params = self._coerce_params(params) param_types_map = self._infer_param_types(coerced_params) if statement.returns_rows(): reader = cast("_SpannerReadProtocol", cursor) execute_kwargs = self._execute_kwargs(for_read=True) result_set = reader.execute_sql(sql, params=coerced_params, param_types=param_types_map, **execute_kwargs) rows = list(result_set) try: metadata = result_set.metadata row_type = metadata.row_type fields = row_type.fields except AttributeError: fields = None if not fields: msg = "Result set metadata not available." raise SQLConversionError(msg) column_names, column_plan = self._resolve_row_plan(fields) data, column_names = collect_rows(rows, fields, column_names=column_names, column_plan=column_plan) return self.create_execution_result( cursor, selected_data=data, column_names=column_names, data_row_count=len(data), is_select_result=True, row_format="tuple", ) if supports_write(cursor): writer = cast("_SpannerWriteProtocol", cursor) execute_kwargs = self._execute_kwargs() row_count = writer.execute_update(sql, params=coerced_params, param_types=param_types_map, **execute_kwargs) return self.create_execution_result(cursor, rowcount_override=row_count) raise SQLConversionError(_READ_ONLY_SNAPSHOT_ERROR_MESSAGE)