"""Oracle Driver"""
import logging
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, overload
from oracledb import create_pipeline as create_oracle_pipeline
from sqlspec.adapters.oracledb._typing import (
DB_TYPE_BLOB,
DB_TYPE_CLOB,
OracleAsyncConnection,
OracleAsyncCursor,
OracleAsyncSessionContext,
OracleSyncConnection,
OracleSyncCursor,
OracleSyncSessionContext,
)
from sqlspec.adapters.oracledb._typing import DatabaseError as OracleDatabaseError
from sqlspec.adapters.oracledb._typing import Error as OracleError
from sqlspec.adapters.oracledb.core import (
ORACLEDB_VERSION,
OracleAsyncStreamSource,
OracleSyncStreamSource,
build_arrow_fetch_kwargs,
build_fetch_kwargs,
build_insert_statement,
build_pipeline_stack_result,
build_truncate_statement,
coerce_large_parameters_async,
coerce_large_parameters_sync,
coerce_many_parameters_async,
coerce_many_parameters_sync,
collect_async_rows,
collect_sync_rows,
connection_is_thin,
create_mapped_exception,
default_statement_config,
driver_profile,
normalize_column_names,
resolve_row_metadata,
resolve_rowcount,
supports_df_batches,
supports_direct_path_load,
)
from sqlspec.adapters.oracledb.data_dictionary import (
OracledbAsyncDataDictionary,
OracledbSyncDataDictionary,
OracleVersionCache,
)
from sqlspec.core import (
SQL,
StackResult,
StatementConfig,
StatementStack,
build_arrow_result_from_table,
create_arrow_result,
get_cache_config,
register_driver_profile,
)
from sqlspec.core.explain import ORACLE_EXPLAIN_PREFIX, ORACLE_MANAGED_EXPLAIN_META_KEY
from sqlspec.driver import (
AsyncDriverAdapterBase,
AsyncRowStream,
BaseAsyncExceptionHandler,
BaseSyncExceptionHandler,
StackExecutionObserver,
SyncDriverAdapterBase,
SyncRowStream,
describe_stack_statement,
hash_stack_operations,
)
from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError, StackExecutionError
from sqlspec.utils.logging import get_logger, log_with_context
from sqlspec.utils.module_loader import ensure_pyarrow
from sqlspec.utils.text import normalize_identifier, quote_identifier, split_qualified_identifier
from sqlspec.utils.type_guards import has_pipeline_capability, is_async_readable
from sqlspec.utils.uuids import uuid4
if TYPE_CHECKING:
from collections.abc import Sequence
from sqlspec.builder import QueryBuilder
from sqlspec.core import ArrowResult, Statement, StatementFilter
from sqlspec.core.stack import StackOperation
from sqlspec.driver import ExecutionResult
from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry
from sqlspec.typing import ArrowRecordBatch, ArrowReturnFormat, ArrowSchema, SchemaT, StatementParameters
__all__ = (
"OracleAsyncDriver",
"OracleAsyncExceptionHandler",
"OracleAsyncSessionContext",
"OracleSyncDriver",
"OracleSyncExceptionHandler",
"OracleSyncSessionContext",
)
logger = get_logger(__name__)
PLAN_STATEMENT_ID_MAX_LENGTH: Final[int] = 30
"""Width of the PLAN_TABLE STATEMENT_ID column."""
PLAN_STATEMENT_ID_PREFIX: Final[str] = "sqlspec_"
PLAN_DISPLAY_SQL: Final[str] = "SELECT plan_table_output FROM TABLE(DBMS_XPLAN.DISPLAY(NULL, :statement_id, 'TYPICAL'))"
PLAN_CLEANUP_SQL: Final[str] = "DELETE FROM plan_table WHERE statement_id = :statement_id"
class OraclePipelineDriver(Protocol):
"""Protocol for Oracle pipeline driver methods used in stack execution."""
statement_config: "StatementConfig"
driver_features: "dict[str, Any]"
def prepare_statement(
self,
statement: "str | Statement | QueryBuilder",
parameters: "tuple[Any, ...] | dict[str, Any] | None",
*,
statement_config: "StatementConfig | None" = None,
kwargs: "dict[str, Any] | None" = None,
) -> "SQL": ...
def _compiled_sql(self, statement: "SQL", statement_config: "StatementConfig") -> "tuple[str, Any]": ...
# Oracle SQL-context byte thresholds (4000 / 2000) live in driver_features so users
# on MAX_STRING_SIZE=EXTENDED databases can override them; defaults are wired in
# core.apply_driver_features and read at the dispatch_execute call sites below.
PIPELINE_MIN_DRIVER_VERSION: "tuple[int, int, int]" = (2, 4, 0)
PIPELINE_MIN_DATABASE_MAJOR: int = 26
class OraclePipelineMixin:
"""Shared helpers for Oracle pipeline execution."""
__slots__ = ()
def _stack_native_blocker(self, stack: "StatementStack") -> "str | None":
for operation in stack.operations:
if operation.method == "execute_arrow":
return "arrow_operation"
if operation.method == "execute_script":
return "script_operation"
return None
def _log_pipeline_skip(self, reason: str, stack: "StatementStack") -> None:
log_level = logging.INFO if reason == "env_override" else logging.DEBUG
log_with_context(
logger,
log_level,
"stack.native_pipeline.skip",
driver=type(self).__name__,
reason=reason,
hashed_operations=hash_stack_operations(stack),
)
def _prepare_pipeline_operation(self, operation: "StackOperation") -> "_CompiledStackOperation":
driver = cast("OraclePipelineDriver", self)
kwargs = dict(operation.keyword_arguments) if operation.keyword_arguments else {}
statement_config = kwargs.pop("statement_config", None)
config = statement_config or driver.statement_config
if operation.method == "execute":
sql_statement = driver.prepare_statement(
operation.statement, operation.arguments, statement_config=config, kwargs=kwargs
)
elif operation.method == "execute_many":
if not operation.arguments:
msg = "execute_many stack operation requires parameter sets"
raise ValueError(msg)
parameter_sets = operation.arguments[0]
filters = operation.arguments[1:]
if isinstance(operation.statement, SQL):
statement_seed = operation.statement.raw_expression or operation.statement.raw_sql
sql_statement = SQL(statement_seed, parameter_sets, statement_config=config, is_many=True, **kwargs)
else:
base_statement = driver.prepare_statement(
operation.statement, filters, statement_config=config, kwargs=kwargs
)
statement_seed = base_statement.raw_expression or base_statement.raw_sql
sql_statement = SQL(statement_seed, parameter_sets, statement_config=config, is_many=True, **kwargs)
else:
msg = f"Unsupported stack operation method: {operation.method}"
raise ValueError(msg)
compiled_sql, prepared_parameters = driver._compiled_sql( # pyright: ignore[reportPrivateUsage]
sql_statement, config
)
summary = describe_stack_statement(operation.statement)
return _CompiledStackOperation(
statement=sql_statement,
sql=compiled_sql,
parameters=prepared_parameters,
method=operation.method,
returns_rows=sql_statement.returns_rows(),
summary=summary,
)
def _add_pipeline_operation(self, pipeline: Any, operation: "_CompiledStackOperation") -> None:
parameters = operation.parameters or []
if operation.method == "execute":
if operation.returns_rows:
pipeline.add_fetchall(operation.sql, parameters)
else:
pipeline.add_execute(operation.sql, parameters)
return
if operation.method == "execute_many":
pipeline.add_executemany(operation.sql, parameters)
return
msg = f"Unsupported pipeline operation: {operation.method}"
raise ValueError(msg)
def _build_stack_results_from_pipeline(
self,
compiled_operations: "Sequence[_CompiledStackOperation]",
pipeline_results: "Sequence[Any]",
continue_on_error: bool,
observer: StackExecutionObserver,
) -> "list[StackResult]":
driver = cast("OraclePipelineDriver", self)
stack_results: list[StackResult] = []
for index, (compiled, result) in enumerate(zip(compiled_operations, pipeline_results, strict=False)):
try:
error = result.error
except AttributeError:
error = None
if error is not None:
stack_error = StackExecutionError(
index,
compiled.summary,
error,
adapter=type(self).__name__,
mode="continue-on-error" if continue_on_error else "fail-fast",
)
if continue_on_error:
observer.record_operation_error(stack_error)
stack_results.append(StackResult.from_error(stack_error))
continue
raise stack_error
stack_results.append(
build_pipeline_stack_result(
compiled.statement,
compiled.method,
compiled.returns_rows,
compiled.parameters,
result,
driver.driver_features,
)
)
return stack_results
def _wrap_pipeline_error(
self, error: Exception, stack: "StatementStack", continue_on_error: bool
) -> StackExecutionError:
mode = "continue-on-error" if continue_on_error else "fail-fast"
return StackExecutionError(
-1, "Oracle pipeline execution failed", error, adapter=type(self).__name__, mode=mode
)
class OracleSyncExceptionHandler(BaseSyncExceptionHandler):
"""Sync Context manager for handling Oracle database exceptions.
Maps Oracle ORA-XXXXX error codes to specific SQLSpec exceptions
for better error handling in application code.
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 issubclass(exc_type, OracleDatabaseError):
self.pending_exception = create_mapped_exception(exc_val)
return True
return False
class OracleAsyncExceptionHandler(BaseAsyncExceptionHandler):
"""Async context manager for handling Oracle database exceptions.
Maps Oracle ORA-XXXXX error codes to specific SQLSpec exceptions
for better error handling in application code.
Uses deferred exception pattern for mypyc compatibility: exceptions
are stored in pending_exception rather than raised from __aexit__
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 issubclass(exc_type, OracleDatabaseError):
self.pending_exception = create_mapped_exception(exc_val)
return True
return False
class OracleSyncDriver(OraclePipelineMixin, SyncDriverAdapterBase):
"""Synchronous Oracle Database driver.
Provides Oracle Database connectivity with parameter style conversion,
error handling, and transaction management.
"""
__slots__ = (
"_data_dictionary",
"_oracle_version_cache",
"_pipeline_support",
"_pipeline_support_reason",
"_row_metadata_cache",
"_transaction_active",
)
dialect = "oracle"
def __init__(
self,
connection: OracleSyncConnection,
statement_config: "StatementConfig | None" = None,
driver_features: "dict[str, Any] | None" = None,
) -> None:
if statement_config is None:
statement_config = default_statement_config.replace(
enable_caching=get_cache_config().compiled_cache_enabled
)
super().__init__(connection=connection, statement_config=statement_config, driver_features=driver_features)
self._data_dictionary: OracledbSyncDataDictionary | None = None
self._pipeline_support: bool | None = None
self._pipeline_support_reason: str | None = None
self._oracle_version_cache: OracleVersionCache | None = None
self._row_metadata_cache: dict[int, tuple[Any, list[str], bool]] = {}
self._transaction_active = False