Source code for sqlspec.adapters.arrow_odbc.driver

"""arrow-odbc sync driver."""

import contextlib
import re
from collections.abc import Iterable, Mapping
from itertools import chain
from typing import TYPE_CHECKING, Any, Final, cast

from sqlspec.adapters.arrow_odbc._typing import ArrowOdbcConnection, ArrowOdbcCursor, ArrowOdbcError, ArrowOdbcRawCursor
from sqlspec.adapters.arrow_odbc.core import (
    build_statement_config,
    create_mapped_exception,
    driver_profile,
    resolve_dialect_from_dbms_name,
)
from sqlspec.adapters.arrow_odbc.data_dictionary import ArrowOdbcDataDictionary
from sqlspec.core import (
    SQL,
    build_arrow_result_from_reader,
    build_arrow_result_from_table,
    get_cache_config,
    register_driver_profile,
)
from sqlspec.driver import BaseSyncExceptionHandler, SyncDriverAdapterBase, SyncRowStream
from sqlspec.driver._common import validate_savepoint_name
from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError
from sqlspec.utils.module_loader import ensure_pyarrow
from sqlspec.utils.text import quote_identifier, split_qualified_identifier

if TYPE_CHECKING:
    from sqlspec.builder import QueryBuilder
    from sqlspec.core import ArrowResult, Statement, StatementConfig, StatementFilter
    from sqlspec.driver import ExecutionResult
    from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry
    from sqlspec.typing import ArrowRecordBatch, ArrowRecordBatchReader, ArrowReturnFormat, StatementParameters

__all__ = ("ArrowOdbcCursor", "ArrowOdbcDriver", "ArrowOdbcExceptionHandler", "resolve_dialect_from_dbms_name")


class ArrowOdbcExceptionHandler(BaseSyncExceptionHandler):
    """Sync context manager handling arrow-odbc exceptions."""

    __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, ArrowOdbcError):
            self.pending_exception = create_mapped_exception(exc_val)
            return True
        return False


class ArrowOdbcStreamSource:
    """Native Arrow ODBC chunk source backed by Arrow record batches."""

    __slots__ = ("_chunk_size", "_driver", "_parameters", "_reader", "_sql")

    def __init__(
        self, driver: "ArrowOdbcDriver", sql: str, parameters: "list[str | None] | None", chunk_size: int
    ) -> None:
        self._driver = driver
        self._sql = sql
        self._parameters = parameters
        self._chunk_size = chunk_size
        self._reader: Any = None

    def start(self) -> None:
        handler = self._driver.handle_database_exceptions()
        with handler:
            reader = self._driver._read_arrow_batches(self._sql, self._parameters, self._chunk_size)
            self._reader = iter(_to_pyarrow_reader(reader))
        self._driver._check_pending_exception(handler)

    def fetch_chunk(self) -> "list[dict[str, Any]]":
        reader = self._reader
        if reader is None:
            return []
        while True:
            try:
                batch = next(reader)
            except StopIteration:
                return []
            rows = cast("list[dict[str, Any]]", batch.to_pylist())
            if rows:
                return rows

    def close(self, error: bool = False) -> None:
        reader = self._reader
        self._reader = None
        close = getattr(reader, "close", None)
        if callable(close):
            with contextlib.suppress(Exception):
                close()


class ArrowOdbcDriver(SyncDriverAdapterBase):
    """Sync driver for generic ODBC connections with Arrow-native transfer."""

    __slots__ = (
        "_chunk_size_val",
        "_data_dictionary",
        "_dbms_name",
        "_dialect",
        "_max_batch_bytes",
        "_max_binary_size_val",
        "_max_text_size_val",
        "_query_timeout_sec_val",
        "_transaction_active",
        "_use_concurrent_fetch",
        "dialect",
    )

    def __init__(
        self,
        connection: "ArrowOdbcConnection",
        statement_config: "StatementConfig | None" = None,
        driver_features: "dict[str, Any] | None" = None,
    ) -> None:
        features = dict(driver_features or {})
        self._dbms_name = self._resolve_dbms_name(connection, features)
        self._dialect = resolve_dialect_from_dbms_name(self._dbms_name)
        statement_dialect = _statement_dialect_for(self._dialect)
        if statement_config is None:
            statement_config = build_statement_config(dialect=statement_dialect).replace(
                enable_caching=get_cache_config().compiled_cache_enabled
            )
        else:
            statement_config = statement_config.replace(dialect=statement_dialect)

        super().__init__(connection=connection, statement_config=statement_config, driver_features=features)
        self._chunk_size_val: int = int(features.get("chunk_size") or 65_536)
        self._max_batch_bytes: int | None = features.get("max_bytes_per_batch")
        self._max_binary_size_val: int | None = features.get("max_binary_size")
        self._max_text_size_val: int | None = features.get("max_text_size")
        self._query_timeout_sec_val: int | None = features.get("query_timeout_sec")
        self._use_concurrent_fetch: bool = bool(features.get("fetch_concurrently", True))
        self.dialect = statement_dialect
        self._data_dictionary: ArrowOdbcDataDictionary | None = None
        self._transaction_active = False
@property def data_dictionary(self) -> "ArrowOdbcDataDictionary": if self._data_dictionary is None: self._data_dictionary = ArrowOdbcDataDictionary(self._dialect) return self._data_dictionary def dispatch_execute(self, cursor: "ArrowOdbcRawCursor", statement: "SQL") -> "ExecutionResult": sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) if self._dialect == "mssql": sql, prepared_parameters = _inline_mssql_pagination_parameters(sql, prepared_parameters) parameters = _odbc_parameters(prepared_parameters) if statement.returns_rows(): reader = self._read_arrow_batches(sql, parameters, self._chunk_size()) table = _reader_to_table(reader) rows = table.to_pylist() column_names = table.column_names return self.create_execution_result( cursor, selected_data=rows, column_names=column_names, data_row_count=table.num_rows, is_select_result=True, row_format="dict", ) cursor.execute(query=sql, parameters=parameters) return self.create_execution_result(cursor, rowcount_override=0)