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