"""AIOSQLite driver implementation for async SQLite operations."""
import asyncio
import random
import sqlite3
from typing import TYPE_CHECKING, Any, cast
import aiosqlite
from sqlspec.adapters.aiosqlite._typing import AiosqliteCursor, AiosqliteRawCursor, AiosqliteSessionContext
from sqlspec.adapters.aiosqlite.core import (
AiosqliteStreamSource,
_execute_and_resolve_metadata,
_execute_fetchall_with_metadata,
build_insert_statement,
collect_rows,
create_mapped_exception,
default_statement_config,
driver_profile,
format_identifier,
normalize_execute_many_parameters,
normalize_execute_parameters,
resolve_rowcount,
run_on_worker_thread,
)
from sqlspec.adapters.aiosqlite.data_dictionary import AiosqliteDataDictionary
from sqlspec.core import ArrowResult, ParameterStyle, TypedParameter, get_cache_config, register_driver_profile
from sqlspec.core.result import DMLResult
from sqlspec.driver import (
AsyncDriverAdapterBase,
AsyncRowStream,
BaseAsyncExceptionHandler,
parameter_value_needs_processing,
type_coercion_fallbacks,
)
from sqlspec.exceptions import SQLSpecError
from sqlspec.utils.type_guards import resolve_row_format
if TYPE_CHECKING:
from collections.abc import Sequence
from sqlspec.adapters.aiosqlite._typing import AiosqliteConnection
from sqlspec.builder import QueryBuilder
from sqlspec.core import SQL, SQLResult, Statement, StatementConfig, StatementFilter
from sqlspec.core.compiler import OperationType
from sqlspec.driver import ExecutionResult
from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry
from sqlspec.typing import StatementParameters
__all__ = (
"AiosqliteCursor",
"AiosqliteDriver",
"AiosqliteExceptionHandler",
"AiosqliteRawCursor",
"AiosqliteSessionContext",
)
class AiosqliteExceptionHandler(BaseAsyncExceptionHandler):
"""Async context manager for handling aiosqlite database exceptions.
Maps SQLite extended result 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:
_ = exc_type
if isinstance(exc_val, (aiosqlite.Error, sqlite3.Error)):
self.pending_exception = create_mapped_exception(exc_val)
return True
return False
[docs]
class AiosqliteDriver(AsyncDriverAdapterBase):
"""AIOSQLite driver for async SQLite database operations."""
__slots__ = ("_data_dictionary", "_rowid_target_cache")
dialect = "sqlite"
[docs]
def __init__(
self,
connection: "AiosqliteConnection",
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: AiosqliteDataDictionary | None = None
self._rowid_target_cache: dict[tuple[str | None, str], bool] = {}
# ─────────────────────────────────────────────────────────────────────────────
# CORE DISPATCH METHODS
# ─────────────────────────────────────────────────────────────────────────────
[docs]
async def dispatch_execute(self, cursor: "AiosqliteRawCursor", statement: "SQL") -> "ExecutionResult":
"""Execute single SQL statement."""
sql, prepared_parameters = self._compiled_sql(statement, self.statement_config)
self._invalidate_rowid_target_cache(statement.operation_type)
normalized_parameters = normalize_execute_parameters(prepared_parameters)
if statement.returns_rows():
fetched_data, description, _affected_rows, last_inserted_id = await run_on_worker_thread(
self.connection,
_execute_fetchall_with_metadata,
self.connection,
sql,
normalized_parameters,
statement.operation_type,
statement.expression,
self._rowid_target_cache,
)
self._invalidate_rowid_target_cache(statement.operation_type)
data, column_names, row_count = collect_rows(fetched_data, description)
row_format = resolve_row_format(data)
return self.create_execution_result(
cursor,
selected_data=data,
column_names=column_names,
data_row_count=row_count,
is_select_result=True,
row_format=row_format,
last_inserted_id=last_inserted_id,
)
affected_rows, last_inserted_id = await run_on_worker_thread(
self.connection,
_execute_and_resolve_metadata,
self.connection,
sql,
normalized_parameters,
statement.operation_type,
statement.expression,
self._rowid_target_cache,
)
self._invalidate_rowid_target_cache(statement.operation_type)
return self.create_execution_result(cursor, rowcount_override=affected_rows, last_inserted_id=last_inserted_id)
[docs]
async def dispatch_execute_many(self, cursor: "AiosqliteRawCursor", statement: "SQL") -> "ExecutionResult":
"""Execute SQL with multiple parameter sets."""
sql, prepared_parameters = self._compiled_sql(statement, self.statement_config)
self._invalidate_rowid_target_cache(statement.operation_type)
try:
await cursor.executemany(sql, normalize_execute_many_parameters(prepared_parameters))
finally:
self._invalidate_rowid_target_cache(statement.operation_type)
affected_rows = resolve_rowcount(cursor)
return self.create_execution_result(cursor, rowcount_override=affected_rows, is_many_result=True)
[docs]
async def dispatch_execute_script(self, cursor: "AiosqliteRawCursor", statement: "SQL") -> "ExecutionResult":
"""Execute SQL script."""
self._rowid_target_cache.clear()
sql, prepared_parameters = self._compiled_sql(statement, self.statement_config)
statements = self.split_script_statements(sql, statement.statement_config, strip_trailing_semicolon=True)
successful_count = 0
last_cursor = cursor
try:
for stmt in statements:
await cursor.execute(stmt, normalize_execute_parameters(prepared_parameters))
successful_count += 1
finally:
self._rowid_target_cache.clear()
return self.create_execution_result(
last_cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True
)
[docs]
async def execute_many(
self,
statement: "SQL | Statement | QueryBuilder",
/,
parameters: "Sequence[StatementParameters]",
*filters: "StatementParameters | StatementFilter",
statement_config: "StatementConfig | None" = None,
**kwargs: Any,
) -> "SQLResult":
"""Execute many with an AIOSQLite thin path for simple qmark batches."""
config = statement_config or self.statement_config
if (
isinstance(statement, str)
and not filters
and not kwargs
and config is self.statement_config
and self.observability.is_idle
and self._can_use_execute_many_thin_path(statement, parameters, config)
):
try:
cursor = await self.connection.executemany(statement, parameters)
except (aiosqlite.Error, sqlite3.Error) as exc:
raise create_mapped_exception(exc) from exc
rowcount = cursor.rowcount
affected_rows = rowcount if isinstance(rowcount, int) and rowcount > 0 else 0
operation = self._resolve_dml_operation_type(statement)
self._invalidate_rowid_target_cache(operation)
return DMLResult(operation, affected_rows)
return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs)
# ─────────────────────────────────────────────────────────────────────────────
# TRANSACTION MANAGEMENT
# ─────────────────────────────────────────────────────────────────────────────
[docs]
async def begin(self) -> None:
"""Begin a database transaction."""
try:
if not self.connection.in_transaction:
await self.connection.execute("BEGIN IMMEDIATE")
except aiosqlite.Error as e:
await _retry_begin_with_backoff(self.connection, e)
[docs]
async def commit(self) -> None:
"""Commit the current transaction."""
try:
await self.connection.commit()
except aiosqlite.Error as e:
msg = f"Failed to commit transaction: {e}"
raise SQLSpecError(msg) from e
[docs]
async def rollback(self) -> None:
"""Rollback the current transaction."""
try:
await self.connection.rollback()
except aiosqlite.Error as e:
msg = f"Failed to rollback transaction: {e}"
raise SQLSpecError(msg) from e
[docs]
def with_cursor(self, connection: "AiosqliteConnection") -> "AiosqliteCursor":
"""Create async context manager for AIOSQLite cursor."""
return AiosqliteCursor(connection)
[docs]
def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "AsyncRowStream[dict[str, Any]] | None":
"""Return a native aiosqlite row stream backed by chunked ``fetchmany``."""
if not statement.returns_rows():
return None
sql, prepared_parameters = self._compiled_sql(statement, self.statement_config)
return AsyncRowStream(AiosqliteStreamSource(self, sql, prepared_parameters, chunk_size))
[docs]
def handle_database_exceptions(self) -> "AiosqliteExceptionHandler":
"""Handle AIOSQLite-specific exceptions."""
return AiosqliteExceptionHandler()
# ─────────────────────────────────────────────────────────────────────────────
# STORAGE API METHODS
# ─────────────────────────────────────────────────────────────────────────────
[docs]
async def select_to_storage(
self,
statement: "SQL | str",
destination: "StorageDestination",
/,
*parameters: Any,
statement_config: "StatementConfig | None" = None,
partitioner: "dict[str, object] | None" = None,
format_hint: "StorageFormat | None" = None,
telemetry: "StorageTelemetry | None" = None,
**kwargs: Any,
) -> "StorageBridgeJob":
"""Execute a query and stream Arrow results into storage."""
self._require_capability("arrow_export_enabled")
arrow_result = await self.select_to_arrow(statement, *parameters, statement_config=statement_config, **kwargs)
async_pipeline = self._storage_pipeline()
telemetry_payload = await self._write_storage_result(
arrow_result, destination, format_hint=format_hint, pipeline=async_pipeline
)
self._attach_partition_telemetry(telemetry_payload, partitioner)
return self._storage_job(telemetry_payload, telemetry)
[docs]
async def load_from_arrow(
self,
table: str,
source: "ArrowResult | Any",
*,
partitioner: "dict[str, object] | None" = None,
overwrite: bool = False,
telemetry: "StorageTelemetry | None" = None,
) -> "StorageBridgeJob":
"""Load Arrow data into SQLite using batched inserts."""
self._require_capability("arrow_import_enabled")
arrow_table = self._coerce_arrow_table(source)
columns, records = self._arrow_table_to_rows(arrow_table)
prepared_records = (
self.prepare_driver_parameters(records, self.statement_config, is_many=True)
if records and self._arrow_rows_need_preparation(arrow_table)
else records
)
owns_transaction = not self.connection.in_transaction
try:
if owns_transaction:
await self.connection.execute("BEGIN IMMEDIATE")
if overwrite:
statement = f"DELETE FROM {format_identifier(table)}"
async with self.with_cursor(self.connection) as cursor:
await cursor.execute(statement)
if records:
insert_sql = build_insert_statement(table, columns)
async with self.with_cursor(self.connection) as cursor:
await cursor.executemany(insert_sql, cast("Any", prepared_records))
if owns_transaction:
await self.connection.commit()
except (aiosqlite.Error, sqlite3.Error) as exc:
if owns_transaction:
await self.connection.rollback()
raise create_mapped_exception(exc) from exc
telemetry_payload = self._ingest_telemetry(arrow_table)
telemetry_payload["destination"] = table
self._attach_partition_telemetry(telemetry_payload, partitioner)
return self._storage_job(telemetry_payload, telemetry)
[docs]
async def load_from_storage(
self,
table: str,
source: "StorageDestination",
*,
file_format: "StorageFormat",
partitioner: "dict[str, object] | None" = None,
overwrite: bool = False,
) -> "StorageBridgeJob":
"""Load staged artifacts from storage into SQLite."""
arrow_table, inbound = await self._read_storage_arrow(source, file_format=file_format)
return await self.load_from_arrow(
table, arrow_table, partitioner=partitioner, overwrite=overwrite, telemetry=inbound
)
# ─────────────────────────────────────────────────────────────────────────────
# UTILITY METHODS
# ─────────────────────────────────────────────────────────────────────────────
@property
def data_dictionary(self) -> "AiosqliteDataDictionary":
"""Get the data dictionary for this driver.
Returns:
Data dictionary instance for metadata queries
"""
if self._data_dictionary is None:
self._data_dictionary = AiosqliteDataDictionary()
return self._data_dictionary
# ─────────────────────────────────────────────────────────────────────────────
# PRIVATE/INTERNAL METHODS
# ─────────────────────────────────────────────────────────────────────────────
[docs]
def collect_rows(self, cursor: Any, fetched: "list[Any]") -> "tuple[list[Any], list[str], int]":
"""Collect aiosqlite rows for the direct execution path."""
return collect_rows(fetched, cursor.description)
[docs]
def resolve_rowcount(self, cursor: Any) -> int:
"""Resolve rowcount from aiosqlite cursor for the direct execution path."""
return resolve_rowcount(cursor)
async def _execute_cache_hit(
self, sql: str, params: "tuple[Any, ...] | list[Any] | dict[str, Any]", cached: Any
) -> "SQLResult":
"""Execute cached queries through the async cursor fast path."""
prepared_params = self.prepare_driver_parameters(params, self.statement_config, prepared_statement=cached)
normalized_parameters = normalize_execute_parameters(prepared_params)
direct_statement: SQL | None = None
exc_handler = self.handle_database_exceptions()
result: SQLResult | None = None
self._invalidate_rowid_target_cache(cached.operation_type)
try:
async with exc_handler:
if cached.operation_profile.returns_rows:
fetched_data, description, _affected_rows, last_inserted_id = await run_on_worker_thread(
self.connection,
_execute_fetchall_with_metadata,
self.connection,
cached.compiled_sql,
normalized_parameters,
cached.operation_type,
cached.processed_state.parsed_expression,
self._rowid_target_cache,
)
data, column_names, row_count = collect_rows(fetched_data, description)
execution_result = self.create_execution_result(
self.connection,
selected_data=data,
column_names=column_names,
data_row_count=row_count,
is_select_result=True,
row_format=resolve_row_format(data),
last_inserted_id=last_inserted_id,
)
direct_statement = self._cached_statement(
sql,
params,
cached,
cast("tuple[Any, ...] | list[Any] | dict[str, Any]", prepared_params),
params_are_simple=True,
compiled_sql=cached.compiled_sql,
)
result = self.build_statement_result(direct_statement, execution_result)
else:
affected_rows, last_inserted_id = await run_on_worker_thread(
self.connection,
_execute_and_resolve_metadata,
self.connection,
cached.compiled_sql,
normalized_parameters,
cached.operation_type,
cached.processed_state.parsed_expression,
self._rowid_target_cache,
)
result = DMLResult(cached.operation_type, affected_rows, last_inserted_id)
self._check_pending_exception(exc_handler)
assert result is not None
return result
finally:
self._invalidate_rowid_target_cache(cached.operation_type)
if direct_statement is not None:
self._release_pooled_statement(direct_statement)
def _invalidate_rowid_target_cache(self, operation_type: "OperationType") -> None:
if operation_type not in {"SELECT", "INSERT", "UPDATE", "DELETE"}:
self._rowid_target_cache.clear()
def _can_use_execute_many_thin_path(
self, statement: str, parameters: "Sequence[StatementParameters]", config: "StatementConfig"
) -> bool:
if type(parameters) is not list:
return False
if not parameters:
return False
if "?" not in statement:
return False
parameter_config = config.parameter_config
if parameter_config.default_parameter_style is not ParameterStyle.QMARK:
return False
if (
parameter_config.default_execution_parameter_style is not None
and parameter_config.default_execution_parameter_style is not ParameterStyle.QMARK
):
return False
if parameter_config.ast_transformer is not None or parameter_config.output_transformer is not None:
return False
if parameter_config.needs_static_script_compilation:
return False
if config.output_transformer is not None or config.statement_transformers:
return False
return self._thin_path_parameters_are_eligible(parameters, parameter_config.type_coercion_map)
@staticmethod
def _thin_path_parameters_are_eligible(
parameters: "list[StatementParameters]", type_coercion_map: "dict[type, Any] | None"
) -> bool:
first_sequence = AiosqliteDriver._as_sequence_parameter_set(parameters[0])
if first_sequence is None:
return False
first_type = type(first_sequence)
row_len = len(first_sequence)
coercion_map = type_coercion_map
has_type_coercion = bool(coercion_map)
fallback_items = type_coercion_fallbacks(coercion_map) if coercion_map else ()
if row_len == 1:
if has_type_coercion and coercion_map is not None:
for param_set in parameters:
sequence = AiosqliteDriver._as_sequence_parameter_set(param_set)
if sequence is None or type(sequence) is not first_type:
return False
if len(sequence) != 1:
return False
if parameter_value_needs_processing(sequence[0], coercion_map, fallback_items):
return False
return True
for param_set in parameters:
sequence = AiosqliteDriver._as_sequence_parameter_set(param_set)
if sequence is None or type(sequence) is not first_type:
return False
if len(sequence) != 1:
return False
if type(sequence[0]) is TypedParameter:
return False
return True
if has_type_coercion and coercion_map is not None:
for param_set in parameters:
sequence = AiosqliteDriver._as_sequence_parameter_set(param_set)
if sequence is None or type(sequence) is not first_type:
return False
if len(sequence) != row_len:
return False
for value in sequence:
if parameter_value_needs_processing(value, coercion_map, fallback_items):
return False
return True
for param_set in parameters:
sequence = AiosqliteDriver._as_sequence_parameter_set(param_set)
if sequence is None or type(sequence) is not first_type:
return False
if len(sequence) != row_len:
return False
for value in sequence:
if type(value) is TypedParameter:
return False
return True
@staticmethod
def _as_sequence_parameter_set(param_set: "StatementParameters") -> "list[Any] | tuple[Any, ...] | None":
if isinstance(param_set, list):
return param_set
if isinstance(param_set, tuple):
return param_set
return None
@staticmethod
def _resolve_dml_operation_type(statement: str) -> "OperationType":
command_keyword = statement.lstrip().split(None, 1)[0].upper() if statement.strip() else "COMMAND"
if command_keyword == "INSERT":
return "INSERT"
if command_keyword == "UPDATE":
return "UPDATE"
if command_keyword == "DELETE":
return "DELETE"
return "COMMAND"
def _connection_in_transaction(self) -> bool:
"""Check if connection is in transaction.
Returns:
True if connection is in an active transaction.
"""
return bool(self.connection.in_transaction)
async def _retry_begin_with_backoff(
connection: "AiosqliteConnection", initial_error: aiosqlite.Error, max_retries: int = 3
) -> None:
"""Retry ``BEGIN IMMEDIATE`` after SQLite reports a busy connection.
Aiosqlite surfaces SQLite lock contention through ``aiosqlite.Error``. Preserve
the existing bounded exponential-backoff behavior for every native error and
report the original failure if all retries are exhausted.
Args:
connection: Aiosqlite connection used to retry the transaction start.
initial_error: Error raised by the first transaction-start attempt.
max_retries: Maximum number of retry attempts.
Raises:
SQLSpecError: If every retry attempt fails.
"""
for attempt in range(max_retries):
delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) # noqa: S311
await asyncio.sleep(delay)
try:
await connection.execute("BEGIN IMMEDIATE")
except aiosqlite.Error:
if attempt == max_retries - 1:
break
else:
return
msg = f"Failed to begin transaction after retries: {initial_error}"
raise SQLSpecError(msg) from initial_error
register_driver_profile("aiosqlite", driver_profile)