From dfa30617a1bfeb60d804ad3156b2870a7decbce2 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Thu, 24 Sep 2026 23:54:56 +0000 Subject: [PATCH 01/12] feat(adapters): optimize SQLite and AioSQLite adapters --- pyproject.toml | 1 + sqlspec/adapters/aiosqlite/_typing.py | 7 +- sqlspec/adapters/aiosqlite/adk/store.py | 4 +- sqlspec/adapters/aiosqlite/config.py | 18 +- sqlspec/adapters/aiosqlite/core.py | 158 ++++++--------- sqlspec/adapters/aiosqlite/data_dictionary.py | 1 - sqlspec/adapters/aiosqlite/driver.py | 90 +++++---- sqlspec/adapters/aiosqlite/events/store.py | 6 +- sqlspec/adapters/aiosqlite/litestar/store.py | 14 +- sqlspec/adapters/aiosqlite/pool.py | 159 +++++++-------- sqlspec/adapters/aiosqlite/type_converter.py | 20 +- sqlspec/adapters/sqlite/adk/store.py | 140 ++++++-------- sqlspec/adapters/sqlite/config.py | 54 +++++- sqlspec/adapters/sqlite/core.py | 165 ++++++---------- sqlspec/adapters/sqlite/data_dictionary.py | 2 - sqlspec/adapters/sqlite/driver.py | 100 +++++----- sqlspec/adapters/sqlite/events/store.py | 6 +- sqlspec/adapters/sqlite/litestar/store.py | 25 +-- sqlspec/adapters/sqlite/pool.py | 181 ++++++++++-------- sqlspec/adapters/sqlite/type_converter.py | 20 +- .../adapters/_shared/_driver_type_system.py | 2 + 21 files changed, 571 insertions(+), 602 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 6158692aa..d9b4b0e80 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -292,6 +292,7 @@ include = [ "sqlspec/adapters/mysql_common.py", # Shared MySQL-family adapter helpers "sqlspec/adapters/**/core.py", # Adapter compiled helpers "sqlspec/adapters/**/type_converter.py", # All adapters type converters + "sqlspec/adapters/sqlite/driver.py", # SQLite synchronous driver "sqlspec/adapters/oracledb/_param_types.py", # Slot-based LOB/JSON parameter wrappers "sqlspec/adapters/oracledb/_json_handlers.py", # Native JSON inputtypehandler / outputtypehandler chain "sqlspec/adapters/oracledb/_uuid_handlers.py", # UUID ↔ RAW(16) inputtypehandler / outputtypehandler chain diff --git a/sqlspec/adapters/aiosqlite/_typing.py b/sqlspec/adapters/aiosqlite/_typing.py index 55b6f9588..711cd0a26 100644 --- a/sqlspec/adapters/aiosqlite/_typing.py +++ b/sqlspec/adapters/aiosqlite/_typing.py @@ -1,4 +1,3 @@ -# pyright: reportCallIssue=false, reportAttributeAccessIssue=false, reportArgumentType=false """AIOSQLite adapter type definitions. This module contains type aliases and classes that are excluded from mypyc @@ -14,8 +13,6 @@ import aiosqlite as aiosqlite_module from typing_extensions import TypeAliasType -_AiosqliteConnection = aiosqlite.Connection - if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType @@ -24,12 +21,12 @@ from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver from sqlspec.core import StatementConfig - AiosqliteConnection: TypeAlias = _AiosqliteConnection + AiosqliteConnection: TypeAlias = aiosqlite.Connection AiosqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] AiosqliteRawCursor: TypeAlias = aiosqlite.Cursor if not TYPE_CHECKING: - AiosqliteConnection = _AiosqliteConnection + AiosqliteConnection = aiosqlite.Connection AiosqliteConnectionFactory = TypeAliasType("AiosqliteConnectionFactory", type[sqlite3.Connection]) AiosqliteRawCursor = aiosqlite.Cursor diff --git a/sqlspec/adapters/aiosqlite/adk/store.py b/sqlspec/adapters/aiosqlite/adk/store.py index f1b94435d..ced1d772b 100644 --- a/sqlspec/adapters/aiosqlite/adk/store.py +++ b/sqlspec/adapters/aiosqlite/adk/store.py @@ -8,7 +8,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 -from sqlspec.adapters.aiosqlite.config import _render_pragmas +from sqlspec.adapters.aiosqlite.config import render_pragmas from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options @@ -1036,7 +1036,7 @@ def _pragma_overrides(config: "AiosqliteConfig") -> "list[tuple[str, str]]": msg = "extension_config['adk']['pragma_overrides'] must be a mapping of PRAGMA names to values" raise ImproperConfigurationError(msg) try: - return _render_pragmas(pragma_overrides) + return render_pragmas(pragma_overrides) except ImproperConfigurationError as exc: msg = str(exc).replace("driver_features['pragmas']", "extension_config['adk']['pragma_overrides']") raise ImproperConfigurationError(msg) from exc diff --git a/sqlspec/adapters/aiosqlite/config.py b/sqlspec/adapters/aiosqlite/config.py index 96ae62afc..2fbd1b382 100644 --- a/sqlspec/adapters/aiosqlite/config.py +++ b/sqlspec/adapters/aiosqlite/config.py @@ -45,6 +45,9 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", + "apply_extension_pragmas", + "extension_pragma_statements", + "render_pragmas", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -440,7 +443,7 @@ async def _close_pool(self) -> None: self.connection_instance = None -def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": +def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": extension_config = cast("dict[str, Any]", config.extension_config) settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) profile = settings.get("pragma_profile", False) @@ -455,7 +458,7 @@ def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" raise ImproperConfigurationError(msg) try: - statements.extend(f"PRAGMA {name} = {value}" for name, value in _render_pragmas(overrides)) + statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) except ImproperConfigurationError as exc: msg = str(exc).replace( "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" @@ -464,12 +467,12 @@ def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str return tuple(statements) -async def _apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: +async def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: for statement in statements: await connection.execute(statement) -def _render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": +def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": rendered: list[tuple[str, str]] = [] for pragma_name, pragma_value in pragmas.items(): if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: @@ -488,6 +491,11 @@ def _render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": return rendered +_extension_pragma_statements = extension_pragma_statements +_apply_extension_pragmas = apply_extension_pragmas +_render_pragmas = render_pragmas + + def _validate_entries(entries: Any, required_keys: "tuple[str, ...]", feature_name: str) -> None: for entry in entries: for required_key in required_keys: @@ -505,7 +513,7 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": return None if "pragmas" in runtime_setup: - runtime_setup["pragmas"] = _render_pragmas(runtime_setup["pragmas"]) + runtime_setup["pragmas"] = render_pragmas(runtime_setup["pragmas"]) row_factory = runtime_setup.get("row_factory") if row_factory is not None and not isinstance(row_factory, str) and not callable(row_factory): diff --git a/sqlspec/adapters/aiosqlite/core.py b/sqlspec/adapters/aiosqlite/core.py index 5141e1321..af1af3297 100644 --- a/sqlspec/adapters/aiosqlite/core.py +++ b/sqlspec/adapters/aiosqlite/core.py @@ -51,8 +51,12 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "execute_and_resolve_metadata", "execute_and_resolve_rowcount", "execute_fetchall_with_description", + "execute_fetchall_with_metadata", + "execute_many_on_worker_thread", + "execute_script_on_worker_thread", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", @@ -83,18 +87,13 @@ SQLITE_READONLY_CODE = 8 SQLITE_DATABASE_LIST_MIN_COLUMNS = 2 SQLITE_TABLE_LIST_MIN_COLUMNS = 5 -SQLITE_TABLE_INFO_MIN_COLUMNS = 2 -SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") async def run_on_worker_thread( connection: "AiosqliteConnection", function: "Callable[..., _T]", *args: Any, **kwargs: Any ) -> _T: """Execute a sqlite3 callable on the aiosqlite worker thread.""" - execute = cast( - "Callable[..., Awaitable[_T]]", - connection._execute, # pyright: ignore[reportPrivateUsage] - ) + execute = cast("Callable[..., Awaitable[_T]]", connection._execute) return await execute(function, *args, **kwargs) @@ -244,15 +243,16 @@ def normalize_execute_parameters(parameters: Any) -> Any: class AiosqliteStreamSource: - """Compiled async chunk source streaming dict rows from an aiosqlite cursor via ``fetchmany``.""" + """Compiled async chunk source streaming dict or tuple rows from an aiosqlite cursor via ``fetchmany``.""" - __slots__ = ("_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") + __slots__ = ("_as_dict", "_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") - def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: + def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int, as_dict: bool = True) -> None: self._driver = driver self._sql = sql self._parameters = parameters self._chunk_size = chunk_size + self._as_dict = as_dict self._cursor: Any = None self._column_names: list[str] | None = None @@ -267,12 +267,16 @@ async def _start(self) -> None: self._cursor = cursor await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - async def fetch_chunk(self) -> "list[dict[str, Any]]": + async def fetch_chunk(self) -> "list[Any]": handler = self._driver.handle_database_exceptions() - rows = await self._driver._run_with_exception_handler(handler, self._cursor.fetchmany, self._chunk_size) + rows: list[Any] = await self._driver._run_with_exception_handler( + handler, self._cursor.fetchmany, self._chunk_size + ) self._driver._check_pending_exception(handler) if not rows: return [] + if not self._as_dict: + return rows if self._column_names is None: self._column_names = [description[0] for description in self._cursor.description] return rows_to_dicts(rows, self._column_names) @@ -368,8 +372,6 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) ): return _create_aiosqlite_error(error, error_code, UniqueViolationError, "unique constraint violation") - # SQLITE_BUSY means another process has the database locked - # SQLITE_LOCKED means another connection has the table/rows locked if error_code == SQLITE_BUSY_CODE or error_name == "SQLITE_BUSY": return _create_aiosqlite_error(error, error_code, DeadlockError, "database busy") if error_code == SQLITE_LOCKED_CODE or error_name == "SQLITE_LOCKED": @@ -377,13 +379,11 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) if "locked" in error_msg or "busy" in error_msg: return _create_aiosqlite_error(error, error_code or 0, DeadlockError, "database locked") - # Query interruption (timeout-like behavior) if error_code == SQLITE_INTERRUPT_CODE or error_name == "SQLITE_INTERRUPT": return _create_aiosqlite_error(error, error_code, OperationCancelledError, "query interrupted") if "interrupt" in error_msg: return _create_aiosqlite_error(error, error_code or 0, OperationCancelledError, "query interrupted") - # Permission errors if error_code == SQLITE_PERM_CODE or error_name == "SQLITE_PERM": return _create_aiosqlite_error(error, error_code, PermissionDeniedError, "permission denied") if error_code == SQLITE_READONLY_CODE or error_name == "SQLITE_READONLY": @@ -439,7 +439,7 @@ def build_profile() -> "DriverParameterProfile": preserve_original_params_for_many=False, json_serializer_strategy="helper", custom_type_coercions={ - bool: _bool_to_int, + bool: int, datetime: _TIME_TO_ISO, date: _TIME_TO_ISO, Decimal: _DECIMAL_TO_STRING, @@ -491,7 +491,7 @@ def _execute_fetchall_with_metadata( eligibility_cache: "dict[tuple[str | None, str], bool]", ) -> "tuple[list[Any], Any, int, int | None]": """Execute a query and return rows plus execution metadata on the worker thread.""" - raw_connection = connection._conn # pyright: ignore[reportPrivateUsage] + raw_connection = connection._conn cursor = cast("Any", raw_connection.execute(sql, normalize_execute_parameters(parameters))) try: fetched_data = cursor.fetchall() @@ -514,7 +514,7 @@ def _execute_and_resolve_metadata( eligibility_cache: "dict[tuple[str | None, str], bool]", ) -> "tuple[int, int | None]": """Execute a statement and resolve rowcount and lastrowid on the worker thread.""" - raw_connection = connection._conn # pyright: ignore[reportPrivateUsage] + raw_connection = connection._conn cursor = raw_connection.execute(sql, normalize_execute_parameters(parameters)) try: rowcount = ( @@ -528,6 +528,44 @@ def _execute_and_resolve_metadata( cast("Any", cursor).close() +def _execute_many_on_worker_thread(connection: "AiosqliteConnection", sql: str, parameters: Any) -> int: + """Execute SQL with multiple parameter sets on the worker thread.""" + raw_connection = connection._conn + cursor = raw_connection.cursor() + try: + cursor.executemany(sql, normalize_execute_many_parameters(parameters)) + return ( + cursor.rowcount if has_rowcount(cursor) and isinstance(cursor.rowcount, int) and cursor.rowcount > 0 else 0 + ) + finally: + with contextlib.suppress(Exception): + cast("Any", cursor).close() + + +def _execute_script_on_worker_thread( + connection: "AiosqliteConnection", statements: "Sequence[str]", parameters: Any +) -> tuple[int, int]: + """Execute multi-statement SQL script on the worker thread.""" + raw_connection = connection._conn + cursor = raw_connection.cursor() + normalized_params = normalize_execute_parameters(parameters) + successful_count = 0 + try: + for stmt in statements: + cursor.execute(stmt, normalized_params) + successful_count += 1 + return len(statements), successful_count + finally: + with contextlib.suppress(Exception): + cursor.close() + + +execute_and_resolve_metadata = _execute_and_resolve_metadata +execute_fetchall_with_metadata = _execute_fetchall_with_metadata +execute_many_on_worker_thread = _execute_many_on_worker_thread +execute_script_on_worker_thread = _execute_script_on_worker_thread + + def _resolve_insert_target(expression: Any) -> "tuple[str | None, str] | None": if not isinstance(expression, exp.Insert): return None @@ -552,7 +590,7 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> with contextlib.suppress(Exception): table_cursor.close() except sqlite3.Error: - return _target_supports_rowid_legacy(connection, target) + return False candidates = [ row @@ -566,10 +604,10 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> if target_schema is not None: candidates = [row for row in candidates if row[0].casefold() == target_schema.casefold()] if not candidates: - return _target_supports_rowid_legacy(connection, target) + return False return len(candidates) == 1 and candidates[0][4] == 0 if not candidates: - return _target_supports_rowid_legacy(connection, target) + return False schema_order = ["temp", "main"] if not any(row[0].casefold() in {"temp", "main"} for row in candidates): @@ -594,80 +632,6 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> return False -def _target_supports_rowid_legacy(connection: Any, target: "tuple[str | None, str]") -> bool: - target_schema, target_table = target - schema_order = [target_schema] if target_schema is not None else ["temp", "main"] - if target_schema is None: - try: - database_cursor = connection.execute("PRAGMA database_list") - try: - schema_order.extend( - row[1] - for row in database_cursor.fetchall() - if len(row) >= SQLITE_DATABASE_LIST_MIN_COLUMNS - and isinstance(row[1], str) - and row[1] not in {"main", "temp"} - ) - finally: - with contextlib.suppress(Exception): - database_cursor.close() - except sqlite3.Error: - return False - - for schema_name in schema_order: - if schema_name is None: - continue - quoted_schema = quote_identifier(schema_name) - schema_cursor = None - schema_row = None - try: - schema_cursor = connection.execute( - f"SELECT type FROM {quoted_schema}.sqlite_master WHERE name = ? COLLATE NOCASE", (target_table,) - ) - schema_row = schema_cursor.fetchone() - except sqlite3.Error: - pass - finally: - if schema_cursor is not None: - with contextlib.suppress(Exception): - schema_cursor.close() - if schema_row is None: - continue - if not schema_row or schema_row[0] != "table": - return False - qualified_target = f"{quoted_schema}.{quote_identifier(target_table)}" - table_info_cursor = None - try: - table_info_cursor = connection.execute( - f"PRAGMA {quoted_schema}.table_info({quote_identifier(target_table)})" - ) - column_names = { - row[1].casefold() - for row in table_info_cursor.fetchall() - if len(row) >= SQLITE_TABLE_INFO_MIN_COLUMNS and isinstance(row[1], str) - } - except sqlite3.Error: - return False - finally: - if table_info_cursor is not None: - with contextlib.suppress(Exception): - table_info_cursor.close() - hidden_alias = next((alias for alias in SQLITE_ROWID_ALIASES if alias not in column_names), None) - if hidden_alias is None: - return False - probe_cursor = None - try: - probe_cursor = connection.execute(f"SELECT {hidden_alias} FROM {qualified_target} LIMIT 0") - except sqlite3.Error: - return False - finally: - if probe_cursor is not None: - with contextlib.suppress(Exception): - probe_cursor.close() - return True - return False - - def _create_aiosqlite_error( error: Any, code: "int | None", error_class: type[SQLSpecError], description: str ) -> SQLSpecError: @@ -689,10 +653,6 @@ def _create_aiosqlite_error( return exc -def _bool_to_int(value: bool) -> int: - return int(value) - - driver_profile = build_profile() default_statement_config = build_statement_config() diff --git a/sqlspec/adapters/aiosqlite/data_dictionary.py b/sqlspec/adapters/aiosqlite/data_dictionary.py index c4fb3778b..97b07fd26 100644 --- a/sqlspec/adapters/aiosqlite/data_dictionary.py +++ b/sqlspec/adapters/aiosqlite/data_dictionary.py @@ -59,7 +59,6 @@ async def get_version(self, driver: "AiosqliteDriver") -> "VersionInfo | None": SQLite version information or None if detection fails. """ driver_id = id(driver) - # Inline cache check to avoid cross-module method call that causes mypyc segfault if driver_id in self._version_fetch_attempted: return self._version_cache.get(driver_id) diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index dcf2d2290..0cf175f56 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -1,7 +1,7 @@ """AIOSQLite driver implementation for async SQLite operations.""" import asyncio -import random +import secrets from typing import TYPE_CHECKING, Any, cast from sqlspec.adapters.aiosqlite._typing import AiosqliteCursor, AiosqliteRawCursor, AiosqliteSessionContext @@ -9,15 +9,16 @@ from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 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, + execute_and_resolve_metadata, + execute_fetchall_with_metadata, + execute_many_on_worker_thread, + execute_script_on_worker_thread, format_identifier, - normalize_execute_many_parameters, normalize_execute_parameters, resolve_rowcount, run_on_worker_thread, @@ -106,7 +107,7 @@ async def dispatch_execute(self, cursor: "AiosqliteRawCursor", statement: "SQL") if statement.returns_rows(): fetched_data, description, _affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - _execute_fetchall_with_metadata, + execute_fetchall_with_metadata, self.connection, sql, normalized_parameters, @@ -130,7 +131,7 @@ async def dispatch_execute(self, cursor: "AiosqliteRawCursor", statement: "SQL") affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - _execute_and_resolve_metadata, + execute_and_resolve_metadata, self.connection, sql, normalized_parameters, @@ -147,12 +148,12 @@ async def dispatch_execute_many(self, cursor: "AiosqliteRawCursor", statement: " self._invalidate_rowid_target_cache(statement.operation_type) try: - await cursor.executemany(sql, normalize_execute_many_parameters(prepared_parameters)) + affected_rows = await run_on_worker_thread( + self.connection, execute_many_on_worker_thread, self.connection, sql, 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) async def dispatch_execute_script(self, cursor: "AiosqliteRawCursor", statement: "SQL") -> "ExecutionResult": @@ -161,18 +162,15 @@ async def dispatch_execute_script(self, cursor: "AiosqliteRawCursor", statement: 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 + statement_count, successful_count = await run_on_worker_thread( + self.connection, execute_script_on_worker_thread, self.connection, statements, prepared_parameters + ) 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 + cursor, statement_count=statement_count, successful_statements=successful_count, is_script_result=True ) async def execute_many( @@ -194,16 +192,19 @@ async def execute_many( and self.observability.is_idle and self._can_use_execute_many_thin_path(statement, parameters, config) ): + cursor = None try: cursor = await self.connection.executemany(statement, parameters) + 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) 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) + finally: + if cursor is not None: + await cursor.close() return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) async def begin(self) -> None: @@ -234,12 +235,14 @@ def with_cursor(self, connection: "AiosqliteConnection") -> "AiosqliteCursor": """Create async context manager for AIOSQLite cursor.""" return AiosqliteCursor(connection) - def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "AsyncRowStream[dict[str, Any]] | None": + def dispatch_select_stream( + self, statement: "SQL", chunk_size: int, as_dict: bool = True + ) -> "AsyncRowStream[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)) + return AsyncRowStream(AiosqliteStreamSource(self, sql, prepared_parameters, chunk_size, as_dict=as_dict)) def handle_database_exceptions(self) -> "AiosqliteExceptionHandler": """Handle AIOSQLite-specific exceptions.""" @@ -273,37 +276,42 @@ async def load_from_arrow( table: str, source: "ArrowResult | Any", *, + batch_size: int = 10000, 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 - ) + columns = arrow_table.column_names + insert_sql = build_insert_statement(table, columns) + needs_prep = self._arrow_rows_need_preparation(arrow_table) + owns_transaction = not self.connection.in_transaction try: if owns_transaction: - await self.connection.execute("BEGIN IMMEDIATE") + await self.begin() 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)) + for batch in arrow_table.to_batches(max_chunksize=batch_size): + pydict = batch.to_pydict() + records = list(zip(*(pydict[col] for col in columns), strict=False)) + if records: + prepared_records = ( + self.prepare_driver_parameters(records, self.statement_config, is_many=True) + if needs_prep + else records + ) + 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() + await self.commit() except (aiosqlite.Error, sqlite3.Error) as exc: if owns_transaction: - await self.connection.rollback() + await self.rollback() raise create_mapped_exception(exc) from exc telemetry_payload = self._ingest_telemetry(arrow_table) @@ -361,7 +369,7 @@ async def _execute_cache_hit( 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, + execute_fetchall_with_metadata, self.connection, cached.compiled_sql, normalized_parameters, @@ -391,7 +399,7 @@ async def _execute_cache_hit( else: affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - _execute_and_resolve_metadata, + execute_and_resolve_metadata, self.connection, cached.compiled_sql, normalized_parameters, @@ -545,7 +553,7 @@ async def _retry_begin_with_backoff( SQLSpecError: If every retry attempt fails. """ for attempt in range(max_retries): - delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) # noqa: S311 + delay = 0.01 * (2**attempt) + secrets.SystemRandom().uniform(0, 0.01) await asyncio.sleep(delay) try: await connection.execute("BEGIN IMMEDIATE") diff --git a/sqlspec/adapters/aiosqlite/events/store.py b/sqlspec/adapters/aiosqlite/events/store.py index e843c78f5..dedff444a 100644 --- a/sqlspec/adapters/aiosqlite/events/store.py +++ b/sqlspec/adapters/aiosqlite/events/store.py @@ -4,7 +4,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.aiosqlite.config import AiosqliteConfig, _apply_extension_pragmas, _extension_pragma_statements +from sqlspec.adapters.aiosqlite.config import AiosqliteConfig, apply_extension_pragmas, extension_pragma_statements from sqlspec.config import EventsConfig from sqlspec.extensions.events import BaseEventQueueStore @@ -36,11 +36,11 @@ class AiosqliteEventQueueStore(BaseEventQueueStore[AiosqliteConfig]): def __init__(self, config: AiosqliteConfig) -> None: super().__init__(config) - self._pragma_statements = _extension_pragma_statements(config, "events") + self._pragma_statements = extension_pragma_statements(config, "events") async def prepare_schema_async(self, driver: Any) -> None: """Apply configured SQLite PRAGMAs before queue DDL.""" - await _apply_extension_pragmas(driver.connection, self._pragma_statements) + await apply_extension_pragmas(driver.connection, self._pragma_statements) def _column_types(self) -> "tuple[str, str, str]": """Return SQLite-compatible column types for the event queue.""" diff --git a/sqlspec/adapters/aiosqlite/litestar/store.py b/sqlspec/adapters/aiosqlite/litestar/store.py index 2d1b1ff03..aece8c9e6 100644 --- a/sqlspec/adapters/aiosqlite/litestar/store.py +++ b/sqlspec/adapters/aiosqlite/litestar/store.py @@ -5,7 +5,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.aiosqlite.config import _apply_extension_pragmas, _extension_pragma_statements +from sqlspec.adapters.aiosqlite.config import apply_extension_pragmas, extension_pragma_statements from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore @@ -59,7 +59,7 @@ def __init__(self, config: "AiosqliteConfig") -> None: config: AiosqliteConfig instance. """ super().__init__(config) - self._pragma_statements = _extension_pragma_statements(config, "litestar") + self._pragma_statements = extension_pragma_statements(config, "litestar") async def create_table(self) -> None: """Create the session table if it doesn't exist.""" @@ -68,14 +68,14 @@ async def create_table(self) -> None: return sql = self._table_ddl() async with self._config.provide_session() as driver: - await _apply_extension_pragmas(driver.connection, self._pragma_statements) + await apply_extension_pragmas(driver.connection, self._pragma_statements) await driver.execute_script(sql) self._log_table_created() await self.reconcile_schema(assume_existing=True) async def prepare_schema_async(self, driver: Any) -> None: """Apply configured SQLite PRAGMAs before migration DDL generation.""" - await _apply_extension_pragmas(driver.connection, self._pragma_statements) + await apply_extension_pragmas(driver.connection, self._pragma_statements) async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Get a session value by key. @@ -90,7 +90,7 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = ? - AND (expires_at IS NULL OR julianday(expires_at) > julianday('now')) + AND (expires_at IS NULL OR expires_at > julianday('now')) """ async with self._config.provide_connection() as conn: @@ -170,7 +170,7 @@ async def exists(self, key: str) -> bool: sql = f""" SELECT 1 FROM {self._table_name} WHERE session_id = ? - AND (expires_at IS NULL OR julianday(expires_at) > julianday('now')) + AND (expires_at IS NULL OR expires_at > julianday('now')) """ async with self._config.provide_connection() as conn, conn.execute(sql, (key,)) as cursor: @@ -218,7 +218,7 @@ async def delete_expired(self) -> int: Returns: Number of sessions deleted. """ - sql = f"DELETE FROM {self._table_name} WHERE julianday(expires_at) <= julianday('now')" + sql = f"DELETE FROM {self._table_name} WHERE expires_at <= julianday('now')" async with self._config.provide_connection() as conn: cursor = await conn.execute(sql) diff --git a/sqlspec/adapters/aiosqlite/pool.py b/sqlspec/adapters/aiosqlite/pool.py index 35695418d..e65594b04 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -16,12 +16,15 @@ from sqlspec.utils.uuids import uuid4 if TYPE_CHECKING: - from collections.abc import Awaitable, Callable + from collections.abc import Awaitable, Callable, Sequence from types import TracebackType from sqlspec.adapters.aiosqlite._typing import AiosqliteConnection __all__ = ( + "SQLITE_DISK_CACHE_SIZE", + "SQLITE_JOURNAL_SIZE_LIMIT", + "SQLITE_MMAP_SIZE", "AiosqliteConnectTimeoutError", "AiosqliteConnectionPool", "AiosqlitePoolClosedError", @@ -36,20 +39,29 @@ SQLITE_DEFAULT_ENABLE_FOREIGN_KEYS: Final = False SQLITE_DEFAULT_ENABLE_OPTIMIZATIONS: Final = True SQLITE_MEMORY_CACHE_SIZE: Final = -16000 +SQLITE_DISK_CACHE_SIZE: Final = -64000 +SQLITE_MMAP_SIZE: Final = 268435456 +SQLITE_JOURNAL_SIZE_LIMIT: Final = 67108864 SQLITE_WAL_SWITCH_ATTEMPTS: Final = 50 SQLITE_WAL_SWITCH_DELAY: Final = 0.01 +async def _attempt_wal_switch(connection: "AiosqliteConnection", attempt: int) -> bool: + """Attempt a single WAL mode switch, returning True on success.""" + try: + await connection.execute("PRAGMA journal_mode = WAL") + except sqlite3.OperationalError as exc: + if "locked" not in str(exc) or attempt == SQLITE_WAL_SWITCH_ATTEMPTS - 1: + raise + await asyncio.sleep(SQLITE_WAL_SWITCH_DELAY) + return False + return True + + async def _enable_wal(connection: "AiosqliteConnection") -> None: """Retry database and table locks briefly while switching to WAL mode.""" for attempt in range(SQLITE_WAL_SWITCH_ATTEMPTS): - try: - await connection.execute("PRAGMA journal_mode = WAL") - except sqlite3.OperationalError as exc: # noqa: PERF203 - bounded lock retry - if "locked" not in str(exc) or attempt == SQLITE_WAL_SWITCH_ATTEMPTS - 1: - raise - await asyncio.sleep(SQLITE_WAL_SWITCH_DELAY) - else: + if await _attempt_wal_switch(connection, attempt): return @@ -71,6 +83,19 @@ def _has_active_transaction(connection: "AiosqliteConnection") -> bool: return bool(getattr(connection, "in_transaction", False)) +def _register_runtime_objects( + connection: "AiosqliteConnection", aggregates: "Sequence[dict[str, Any]]", collations: "Sequence[dict[str, Any]]" +) -> None: + """Register custom aggregates and collations on the worker thread.""" + raw_connection = connection._conn + for aggregate_config in aggregates: + raw_connection.create_aggregate( + aggregate_config["name"], aggregate_config["narg"], aggregate_config["aggregate_class"] + ) + for collation_config in collations: + raw_connection.create_collation(collation_config["name"], collation_config["func"]) + + async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: "dict[str, Any]") -> None: pragmas = runtime_setup.get("pragmas", ()) if pragmas: @@ -94,20 +119,10 @@ async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: deterministic=function_config.get("deterministic", False), ) - raw_connection = connection._conn # pyright: ignore[reportPrivateUsage] - for aggregate_config in runtime_setup.get("custom_aggregates", ()): - await run_on_worker_thread( - connection, - raw_connection.create_aggregate, - aggregate_config["name"], - aggregate_config["narg"], - aggregate_config["aggregate_class"], - ) - - for collation_config in runtime_setup.get("custom_collations", ()): - await run_on_worker_thread( - connection, raw_connection.create_collation, collation_config["name"], collation_config["func"] - ) + aggregates = runtime_setup.get("custom_aggregates", ()) + collations = runtime_setup.get("custom_collations", ()) + if aggregates or collations: + await run_on_worker_thread(connection, _register_runtime_objects, connection, aggregates, collations) authorizer_callback = runtime_setup.get("authorizer_callback") if authorizer_callback is not None: @@ -231,7 +246,6 @@ async def close(self) -> None: await self.connection.rollback() await self.connection.close() except Exception: - # Note: No pool context available at connection level log_with_context( logger, logging.DEBUG, "pool.connection.close.error", adapter=_ADAPTER_NAME, connection_id=self.id ) @@ -271,7 +285,7 @@ class AiosqliteConnectionPool: """Multi-connection pool for aiosqlite.""" __slots__ = ( - "_closed_event_instance", + "_closed_event", "_connect_timeout", "_connection_parameters", "_connection_registry", @@ -279,13 +293,13 @@ class AiosqliteConnectionPool: "_enable_optimizations", "_health_check_interval", "_idle_timeout", - "_lock_instance", + "_lock", "_min_size", "_on_connection_create", "_operation_timeout", "_pool_id", "_pool_size", - "_queue_instance", + "_queue", "_runtime_setup", "_warmed", ) @@ -333,32 +347,11 @@ def __init__( self._connection_registry: dict[str, AiosqlitePoolConnection] = {} self._warmed = False - self._pool_id = uuid4().hex[:8] # Short ID for logging - - self._queue_instance: asyncio.Queue[AiosqlitePoolConnection] | None = None - self._lock_instance: asyncio.Lock | None = None - self._closed_event_instance: asyncio.Event | None = None + self._pool_id = uuid4().hex[:8] - @property - def _queue(self) -> "asyncio.Queue[AiosqlitePoolConnection]": - """Lazy initialization of asyncio.Queue for Python 3.9 compatibility.""" - if self._queue_instance is None: - self._queue_instance = asyncio.Queue(maxsize=self._pool_size) - return self._queue_instance - - @property - def _lock(self) -> asyncio.Lock: - """Lazy initialization of asyncio.Lock for Python 3.9 compatibility.""" - if self._lock_instance is None: - self._lock_instance = asyncio.Lock() - return self._lock_instance - - @property - def _closed_event(self) -> asyncio.Event: - """Lazy initialization of asyncio.Event for Python 3.9 compatibility.""" - if self._closed_event_instance is None: - self._closed_event_instance = asyncio.Event() - return self._closed_event_instance + self._queue: asyncio.Queue[AiosqlitePoolConnection] = asyncio.Queue(maxsize=self._pool_size) + self._lock: asyncio.Lock = asyncio.Lock() + self._closed_event: asyncio.Event = asyncio.Event() @property def is_closed(self) -> bool: @@ -367,7 +360,7 @@ def is_closed(self) -> bool: Returns: True if pool is closed """ - return self._closed_event_instance is not None and self._closed_event.is_set() + return self._closed_event.is_set() @property def _database_name(self) -> str: @@ -376,17 +369,13 @@ def _database_name(self) -> str: return str(db).split("/")[-1] if db else "unknown" def _set_connect_proxy_daemon(self, connect_proxy: Any) -> None: - """Set daemon mode on aiosqlite worker thread before await. - - aiosqlite <=0.21 used Connection as a Thread subclass. - aiosqlite >=0.22 stores an internal ``_thread`` attribute instead. - """ + """Set daemon mode on aiosqlite worker thread before await.""" try: if isinstance(connect_proxy, Thread): connect_proxy.daemon = True return - worker_thread = connect_proxy._thread # pyright: ignore[reportAttributeAccessIssue] + worker_thread = getattr(connect_proxy, "_thread", None) if isinstance(worker_thread, Thread): worker_thread.daemon = True except Exception: @@ -401,8 +390,23 @@ def _set_connect_proxy_daemon(self, connect_proxy: Any) -> None: async def _force_stop_connection(self, connection: AiosqlitePoolConnection, *, reason: str) -> None: """Force-stop aiosqlite worker thread when graceful close times out.""" + raw_conn = getattr(connection.connection, "_conn", None) + if raw_conn is not None: + with suppress(Exception): + raw_conn.interrupt() try: - stop_method = connection.connection.stop # pyright: ignore[reportAttributeAccessIssue] + stop_method = getattr(connection.connection, "stop", None) + if stop_method is None: + log_with_context( + logger, + logging.DEBUG, + "pool.connection.force_stop.unavailable", + adapter=_ADAPTER_NAME, + pool_id=self._pool_id, + connection_id=connection.id, + reason=reason, + ) + return except Exception: log_with_context( logger, @@ -465,8 +469,6 @@ def checked_out(self) -> int: Returns: Number of connections currently in use """ - if self._queue_instance is None: - return len(self._connection_registry) return len(self._connection_registry) - self._queue.qsize() async def _create_connection(self) -> AiosqlitePoolConnection: @@ -517,8 +519,19 @@ async def new_connection(self) -> "AiosqliteConnection": f"PRAGMA cache_size = {SQLITE_MEMORY_CACHE_SIZE}", ]) else: - await _enable_wal(connection) - pragma_lines.append("PRAGMA synchronous = NORMAL") + cursor = await connection.execute("PRAGMA journal_mode") + current_mode = await cursor.fetchone() + await cursor.close() + if not current_mode or str(current_mode[0]).upper() != "WAL": + await _enable_wal(connection) + pragma_lines.extend([ + "PRAGMA synchronous = NORMAL", + "PRAGMA temp_store = MEMORY", + f"PRAGMA mmap_size = {SQLITE_MMAP_SIZE}", + f"PRAGMA cache_size = {SQLITE_DISK_CACHE_SIZE}", + f"PRAGMA journal_size_limit = {SQLITE_JOURNAL_SIZE_LIMIT}", + "PRAGMA threads = 4", + ]) pragma_lines.append(f"PRAGMA busy_timeout = {SQLITE_BUSY_TIMEOUT}") @@ -645,8 +658,6 @@ async def _try_provision_new_connection(self) -> "AiosqlitePoolConnection | None try: connection = await self._create_connection() except Exception: - # Surface the real cause (bad on_connection_create hook, bad DSN, disk full) instead of - # returning None and letting acquire() stall on an empty queue until connect_timeout. log_with_context( logger, logging.WARNING, @@ -744,29 +755,23 @@ async def _get_connection(self) -> AiosqlitePoolConnection: Raises: AiosqlitePoolClosedError: If pool is closed """ - # Fast path: check closed state directly to avoid property overhead - if self._closed_event_instance is not None and self._closed_event_instance.is_set(): + if self._closed_event.is_set(): msg = "Cannot acquire connection from closed pool" raise AiosqlitePoolClosedError(msg) if not self._warmed and self._min_size > 0: await self._warm_pool() - # Fast path: try to get from queue without health check overhead for fresh connections while not self._queue.empty(): connection = self._queue.get_nowait() - # Fast claim for recently-used connections (idle < health_check_interval) if connection.idle_since is not None: idle_time = time.time() - connection.idle_since if idle_time <= self._health_check_interval and connection.is_healthy: connection.idle_since = None return connection - # Fall back to full health check for older connections if await self._claim_if_healthy(connection): return connection - # Try to create new connection if under capacity - # Fast path: check capacity without lock first if len(self._connection_registry) < self._pool_size: new_connection = await self._try_provision_new_connection() if new_connection is not None: @@ -799,8 +804,7 @@ async def release(self, connection: AiosqlitePoolConnection) -> None: Args: connection: Connection to release """ - # Fast path: check closed state directly - if self._closed_event_instance is not None and self._closed_event_instance.is_set(): + if self._closed_event.is_set(): await self._retire_connection(connection) return @@ -816,8 +820,6 @@ async def release(self, connection: AiosqlitePoolConnection) -> None: return try: - # Fast path: skip timeout wrapper for reset, just do the rollback directly - # The rollback itself is fast for SQLite; timeout is overkill for hot path if _has_active_transaction(connection.connection): with suppress(Exception): await connection.connection.rollback() @@ -854,6 +856,11 @@ async def close(self) -> None: self._connection_registry.clear() if connections: + for conn in connections: + raw_conn = getattr(conn.connection, "_conn", None) + if raw_conn is not None: + with suppress(Exception): + raw_conn.interrupt() close_tasks = [asyncio.wait_for(conn.close(), timeout=self._operation_timeout) for conn in connections] results = await asyncio.gather(*close_tasks, return_exceptions=True) diff --git a/sqlspec/adapters/aiosqlite/type_converter.py b/sqlspec/adapters/aiosqlite/type_converter.py index d4fa27f7b..3579d1d7c 100644 --- a/sqlspec/adapters/aiosqlite/type_converter.py +++ b/sqlspec/adapters/aiosqlite/type_converter.py @@ -1,4 +1,3 @@ -# Keep in sync with sqlspec/adapters/sqlite/type_converter.py """SQLite custom type handlers for optional JSON and type conversion support. Provides registration functions for SQLite's adapter/converter system to enable @@ -9,12 +8,12 @@ instead of lambdas for adapter registration. """ -import json from functools import partial from typing import TYPE_CHECKING, Any from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 from sqlspec.utils.logging import get_logger +from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Callable @@ -31,14 +30,13 @@ def json_adapter(value: Any, serializer: "Callable[[Any], str] | None" = None) - Args: value: Python dict or list to serialize. - serializer: Optional JSON serializer callable. Defaults to standard json.dumps. + serializer: Optional JSON serializer callable. Defaults to compiled to_json. Returns: JSON string representation. """ - if serializer is None: - return json.dumps(value, ensure_ascii=False) - return serializer(value) + codec = serializer or to_json + return codec(value) def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = None) -> Any: @@ -46,14 +44,13 @@ def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = N Args: value: UTF-8 encoded JSON bytes from SQLite. - deserializer: Optional JSON deserializer callable. Defaults to standard json.loads. + deserializer: Optional JSON deserializer callable. Defaults to compiled from_json. Returns: Deserialized Python object (dict or list). """ - if deserializer is None: - return json.loads(value.decode("utf-8")) - return deserializer(value.decode("utf-8")) + codec = deserializer or from_json + return codec(value.decode("utf-8")) def register_type_handlers( @@ -64,6 +61,9 @@ def register_type_handlers( This function registers handlers globally for the sqlite3 module. It should be called once during application initialization if custom type handling is needed. + Note that sqlite3.register_adapter is deprecated in Python 3.12+ in favor of + statement-level conversions. + Args: json_serializer: Optional custom JSON serializer. json_deserializer: Optional custom JSON deserializer. diff --git a/sqlspec/adapters/sqlite/adk/store.py b/sqlspec/adapters/sqlite/adk/store.py index 9297a883d..debb39e2e 100644 --- a/sqlspec/adapters/sqlite/adk/store.py +++ b/sqlspec/adapters/sqlite/adk/store.py @@ -8,7 +8,8 @@ from typing_extensions import NotRequired from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 -from sqlspec.adapters.sqlite.config import _render_pragmas +from sqlspec.adapters.sqlite.config import render_pragmas +from sqlspec.adapters.sqlite.core import end_transaction from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options @@ -133,9 +134,8 @@ def create_session( params = (session_id, app_name, user_id, state_json, now_julian, now_julian) with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, params) - conn.commit() + end_transaction(conn, commit=True) return StoredSession( id=session_id, app_name=app_name, user_id=user_id, state=state, create_time=now, update_time=now @@ -155,7 +155,6 @@ def get_session( Returns: Session record or None if not found. """ - """Synchronous implementation of get_session.""" params = (app_name, user_id, session_id) update_params: tuple[Any, ...] if renew_for is not None and self._calculate_expires_at(renew_for) is not None: @@ -178,10 +177,9 @@ def get_session( try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) if update_sql: conn.execute(update_sql, update_params) - conn.commit() + end_transaction(conn, commit=True) cursor = conn.execute(sql, params) row = cursor.fetchone() @@ -210,7 +208,6 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta session_id: Session identifier. state: New state dictionary (replaces existing state). """ - """Synchronous implementation of update_session_state.""" now_julian = _datetime_to_julian(datetime.now(timezone.utc)) state_json = to_json(state) @@ -221,9 +218,8 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, (state_json, now_julian, app_name, user_id, session_id)) - conn.commit() + end_transaction(conn, commit=True) def list_sessions( self, @@ -258,7 +254,6 @@ def list_sessions( try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, params) rows = cursor.fetchall() @@ -286,13 +281,11 @@ def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: user_id: User identifier. session_id: Session identifier. """ - """Synchronous implementation of delete_session.""" sql = f"DELETE FROM {self._session_table} WHERE app_name = ? AND user_id = ? AND id = ?" with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, (app_name, user_id, session_id)) - conn.commit() + end_transaction(conn, commit=True) def append_event(self, event_record: StoredEvent) -> None: """Append an event to a session. @@ -300,7 +293,6 @@ def append_event(self, event_record: StoredEvent) -> None: Args: event_record: Event record to store. """ - """Synchronous implementation of append_event.""" timestamp_julian = _datetime_to_julian(event_record["timestamp"]) event_data_json = to_json(event_record["event_data"]) @@ -311,7 +303,6 @@ def append_event(self, event_record: StoredEvent) -> None: """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute( sql, ( @@ -324,7 +315,7 @@ def append_event(self, event_record: StoredEvent) -> None: event_data_json, ), ) - conn.commit() + end_transaction(conn, commit=True) def append_event_and_update_state( self, @@ -352,7 +343,6 @@ def append_event_and_update_state( app_state: App-scoped state snapshot to upsert when changed. user_state: User-scoped state snapshot to upsert when changed. """ - """Synchronous implementation of append_event_and_update_state.""" timestamp_julian = _datetime_to_julian(event_record["timestamp"]) event_data_json = to_json(event_record["event_data"]) now_julian = _datetime_to_julian(datetime.now(timezone.utc)) @@ -388,7 +378,6 @@ def append_event_and_update_state( """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) try: cursor = conn.execute(update_sql, (state_json, now_julian, app_name, user_id, session_id)) row = cursor.fetchone() @@ -410,13 +399,13 @@ def append_event_and_update_state( if user_state is not None: conn.execute(user_upsert_sql, (app_name, user_id, to_json(user_state), now_julian)) except Exception: - conn.rollback() + end_transaction(conn, commit=False) raise else: if row is None: - conn.rollback() + end_transaction(conn, commit=False) else: - conn.commit() + end_transaction(conn, commit=True) if row is None: msg = f"Session {session_id} not found during append_event_and_update_state." @@ -474,7 +463,6 @@ def get_events( try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, params) rows = cursor.fetchall() @@ -497,7 +485,6 @@ def get_events( def delete_expired_events(self, before: datetime, app_name: "str | None" = None) -> int: """Delete events older than the given timestamp.""" - """Synchronous implementation of delete_expired_events.""" sql = f"DELETE FROM {self._events_table} WHERE timestamp < ?" params: list[Any] = [_datetime_to_julian(before)] if app_name is not None: @@ -506,10 +493,9 @@ def delete_expired_events(self, before: datetime, app_name: "str | None" = None) try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - conn.commit() + end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -526,10 +512,9 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - conn.commit() + end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -538,7 +523,6 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" def delete_idle_user_states(self, updated_before: datetime, app_name: "str | None" = None) -> int: """Delete user state rows whose update_time predates the given threshold.""" - """Synchronous implementation of delete_idle_user_states.""" sql = f"DELETE FROM {self._user_state_table} WHERE update_time < ?" params: list[Any] = [_datetime_to_julian(updated_before)] if app_name is not None: @@ -547,10 +531,9 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - conn.commit() + end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -559,12 +542,10 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non def get_app_state(self, app_name: str) -> "dict[str, Any] | None": """Return app-scoped state for an application.""" - """Synchronous implementation of get_app_state.""" sql = f"SELECT state FROM {self._app_state_table} WHERE app_name = ?" try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, (app_name,)) row = cursor.fetchone() return from_json(row[0]) if row is not None and row[0] else None @@ -575,7 +556,6 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None": """Return user-scoped state for an application user.""" - """Synchronous implementation of get_user_state.""" sql = f""" SELECT state FROM {self._user_state_table} @@ -584,7 +564,6 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, (app_name, user_id)) row = cursor.fetchone() return from_json(row[0]) if row is not None and row[0] else None @@ -595,7 +574,6 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """Insert or replace app-scoped state for an application.""" - """Synchronous implementation of upsert_app_state.""" sql = f""" INSERT INTO {self._app_state_table} (app_name, state, update_time) VALUES (?, ?, ?) @@ -605,13 +583,11 @@ def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, (app_name, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))) - conn.commit() + end_transaction(conn, commit=True) def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: """Insert or replace user-scoped state for an application user.""" - """Synchronous implementation of upsert_user_state.""" sql = f""" INSERT INTO {self._user_state_table} (app_name, user_id, state, update_time) VALUES (?, ?, ?, ?) @@ -621,18 +597,15 @@ def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]" """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, (app_name, user_id, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))) - conn.commit() + end_transaction(conn, commit=True) def get_metadata(self, key: str) -> "str | None": """Return a value from the ADK internal metadata table.""" - """Synchronous implementation of get_metadata.""" sql = f"SELECT value FROM {self._metadata_table} WHERE key = ?" try: with self._config.provide_connection() as conn: - self._apply_pragmas(conn) cursor = conn.execute(sql, (key,)) row = cursor.fetchone() return str(row[0]) if row is not None else None @@ -643,7 +616,6 @@ def get_metadata(self, key: str) -> "str | None": def set_metadata(self, key: str, value: str) -> None: """Set a value in the ADK internal metadata table.""" - """Synchronous implementation of set_metadata.""" sql = f""" INSERT INTO {self._metadata_table} (key, value) VALUES (?, ?) @@ -651,9 +623,8 @@ def set_metadata(self, key: str, value: str) -> None: """ with self._config.provide_connection() as conn: - self._apply_pragmas(conn) conn.execute(sql, (key, value)) - conn.commit() + end_transaction(conn, commit=True) def _apply_pragmas(self, connection: Any) -> None: """Apply PRAGMA optimization profile for this connection. @@ -661,10 +632,6 @@ def _apply_pragmas(self, connection: Any) -> None: Args: connection: SQLite connection. """ - connection.execute("PRAGMA foreign_keys = ON") - connection.execute("PRAGMA cache_size = -64000") - connection.execute("PRAGMA mmap_size = 30000000") - connection.execute("PRAGMA journal_size_limit = 67108864") for pragma_name, pragma_value in self._pragma_overrides: connection.execute(f"PRAGMA {pragma_name} = {pragma_value}") @@ -838,26 +805,25 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object if not entries: return 0 - inserted_count = 0 with self._config.provide_connection() as conn: self._enable_foreign_keys(conn) - for entry in entries: - timestamp_julian = _datetime_to_julian(entry["timestamp"]) - inserted_at_julian = _datetime_to_julian(entry["inserted_at"]) - content_json_str = to_json(entry["content_json"]) - metadata_json_str = to_json(entry["metadata_json"]) if entry["metadata_json"] else None - scope = entry.get("scope", "user") - - if self._owner_id_column_name: - sql = f""" - INSERT OR IGNORE INTO {self._memory_table} - (id, session_id, app_name, user_id, scope, event_id, author, - {self._owner_id_column_name}, timestamp, content_json, - content_text, metadata_json, inserted_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - """ - params: tuple[Any, ...] = ( + params_list: list[tuple[Any, ...]] = [] + if self._owner_id_column_name: + sql = f""" + INSERT OR IGNORE INTO {self._memory_table} + (id, session_id, app_name, user_id, scope, event_id, author, + {self._owner_id_column_name}, timestamp, content_json, + content_text, metadata_json, inserted_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """ + for entry in entries: + timestamp_julian = _datetime_to_julian(entry["timestamp"]) + inserted_at_julian = _datetime_to_julian(entry["inserted_at"]) + content_json_str = to_json(entry["content_json"]) + metadata_json_str = to_json(entry["metadata_json"]) if entry["metadata_json"] else None + scope = entry.get("scope", "user") + params_list.append(( entry["id"], entry["session_id"], entry["app_name"], @@ -871,15 +837,21 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["content_text"], metadata_json_str, inserted_at_julian, - ) - else: - sql = f""" - INSERT OR IGNORE INTO {self._memory_table} - (id, session_id, app_name, user_id, scope, event_id, author, - timestamp, content_json, content_text, metadata_json, inserted_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - """ - params = ( + )) + else: + sql = f""" + INSERT OR IGNORE INTO {self._memory_table} + (id, session_id, app_name, user_id, scope, event_id, author, + timestamp, content_json, content_text, metadata_json, inserted_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """ + for entry in entries: + timestamp_julian = _datetime_to_julian(entry["timestamp"]) + inserted_at_julian = _datetime_to_julian(entry["inserted_at"]) + content_json_str = to_json(entry["content_json"]) + metadata_json_str = to_json(entry["metadata_json"]) if entry["metadata_json"] else None + scope = entry.get("scope", "user") + params_list.append(( entry["id"], entry["session_id"], entry["app_name"], @@ -892,13 +864,11 @@ def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object entry["content_text"], metadata_json_str, inserted_at_julian, - ) - - cursor = conn.execute(sql, params) - if cursor.rowcount > 0: - inserted_count += 1 + )) - conn.commit() + cursor = conn.executemany(sql, params_list) + inserted_count = cursor.rowcount if cursor.rowcount >= 0 else len(params_list) + end_transaction(conn, commit=True) return inserted_count @@ -921,7 +891,7 @@ def search_entries( if self._use_fts: try: return self._search_entries_fts(query, app_name, user_id, effective_limit, scope_filter) - except Exception as exc: # pragma: no cover + except Exception as exc: logger.warning("FTS search failed; falling back to simple search: %s", exc) return self._search_entries_simple(query, app_name, user_id, effective_limit, scope_filter) @@ -933,7 +903,7 @@ def delete_entries_by_session(self, session_id: str) -> int: self._enable_foreign_keys(conn) cursor = conn.execute(sql, (session_id,)) deleted_count = cursor.rowcount - conn.commit() + end_transaction(conn, commit=True) return deleted_count @@ -954,7 +924,7 @@ def delete_entries_older_than(self, days: int, app_name: "str | None" = None, sc self._enable_foreign_keys(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount - conn.commit() + end_transaction(conn, commit=True) return deleted_count @@ -1116,7 +1086,7 @@ def _pragma_overrides(config: "SqliteConfig") -> "list[tuple[str, str]]": msg = "extension_config['adk']['pragma_overrides'] must be a mapping of PRAGMA names to values" raise ImproperConfigurationError(msg) try: - return _render_pragmas(pragma_overrides) + return render_pragmas(pragma_overrides) except ImproperConfigurationError as exc: msg = str(exc).replace("driver_features['pragmas']", "extension_config['adk']['pragma_overrides']") raise ImproperConfigurationError(msg) from exc diff --git a/sqlspec/adapters/sqlite/config.py b/sqlspec/adapters/sqlite/config.py index f000e46c5..b529653d7 100644 --- a/sqlspec/adapters/sqlite/config.py +++ b/sqlspec/adapters/sqlite/config.py @@ -37,6 +37,10 @@ "SqliteConnectionParams", "SqliteDriverFeatures", "SqliteFunctionConfig", + "SqliteWindowFunctionConfig", + "apply_extension_pragmas", + "extension_pragma_statements", + "render_pragmas", ) logger = get_logger("sqlspec.adapters.sqlite") @@ -58,6 +62,10 @@ class SqliteConnectionParams(TypedDict): health_check_interval: NotRequired[float] enable_optimizations: NotRequired[bool] enable_foreign_keys: NotRequired[bool] + busy_timeout: NotRequired[int] + cache_size: NotRequired[int] + mmap_size: NotRequired[int] + default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] extra: NotRequired[dict[str, Any]] @@ -85,6 +93,14 @@ class SqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" +class SqliteWindowFunctionConfig(TypedDict): + """User-defined SQLite window function registration.""" + + name: str + narg: int + window_class: "type[Any]" + + class SqliteDriverFeatures(TypedDict): """SQLite driver feature configuration. @@ -114,6 +130,9 @@ class SqliteDriverFeatures(TypedDict): Each entry must include name and func. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. + custom_window_functions: Register user-defined aggregate window functions. + Each entry must include name, narg, and window_class. + default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation. trace_callback: sqlite3 trace hook run for executed statements. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -137,6 +156,8 @@ class SqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[SqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[SqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[SqliteAggregateConfig]]" + custom_window_functions: "NotRequired[Sequence[SqliteWindowFunctionConfig]]" + default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -155,6 +176,7 @@ class SqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -251,7 +273,6 @@ def __init__( statement_config = statement_config or default_statement_config statement_config, driver_features = apply_driver_features(statement_config, driver_features) - # Extract user connection hook before storing driver_features features_dict = dict(driver_features) if driver_features else {} self._user_connection_hook: Callable[[SqliteConnection], None] | None = features_dict.pop( "on_connection_create", None @@ -305,6 +326,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "SqliteExceptionHandler": SqliteExceptionHandler, "SqliteFunctionConfig": SqliteFunctionConfig, "SqliteSessionContext": SqliteSessionContext, + "SqliteWindowFunctionConfig": SqliteWindowFunctionConfig, }) return namespace @@ -329,6 +351,18 @@ def _create_pool(self) -> SqliteConnectionPool: if enable_foreign_keys is not None: pool_kwargs["enable_foreign_keys"] = enable_foreign_keys + busy_timeout = self.connection_config.get("busy_timeout") + if busy_timeout is not None: + pool_kwargs["busy_timeout"] = busy_timeout + + cache_size = self.connection_config.get("cache_size") + if cache_size is not None: + pool_kwargs["cache_size"] = cache_size + + mmap_size = self.connection_config.get("mmap_size") + if mmap_size is not None: + pool_kwargs["mmap_size"] = mmap_size + pool = SqliteConnectionPool( connection_parameters=config_dict, on_connection_create=self._user_connection_hook, @@ -358,7 +392,7 @@ def _close_pool(self) -> None: self.connection_instance.close() -def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": +def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": extension_config = cast("dict[str, Any]", config.extension_config) settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) profile = settings.get("pragma_profile", False) @@ -373,7 +407,7 @@ def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" raise ImproperConfigurationError(msg) try: - statements.extend(f"PRAGMA {name} = {value}" for name, value in _render_pragmas(overrides)) + statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) except ImproperConfigurationError as exc: msg = str(exc).replace( "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" @@ -382,12 +416,12 @@ def _extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str return tuple(statements) -def _apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: +def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: for statement in statements: connection.execute(statement) -def _render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": +def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": rendered: list[tuple[str, str]] = [] for pragma_name, pragma_value in pragmas.items(): if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: @@ -406,6 +440,11 @@ def _render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": return rendered +_extension_pragma_statements = extension_pragma_statements +_apply_extension_pragmas = apply_extension_pragmas +_render_pragmas = render_pragmas + + def _validate_entries(entries: Any, required_keys: "tuple[str, ...]", feature_name: str) -> None: for entry in entries: for required_key in required_keys: @@ -423,7 +462,7 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": return None if "pragmas" in runtime_setup: - runtime_setup["pragmas"] = _render_pragmas(runtime_setup["pragmas"]) + runtime_setup["pragmas"] = render_pragmas(runtime_setup["pragmas"]) row_factory = runtime_setup.get("row_factory") if row_factory is not None and not isinstance(row_factory, str) and not callable(row_factory): @@ -438,6 +477,9 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) + _validate_entries( + runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" + ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/sqlite/core.py b/sqlspec/adapters/sqlite/core.py index 39a7e56e9..af4a02a4e 100644 --- a/sqlspec/adapters/sqlite/core.py +++ b/sqlspec/adapters/sqlite/core.py @@ -36,9 +36,11 @@ if TYPE_CHECKING: from collections.abc import Callable, Sequence + from sqlspec.adapters.sqlite._typing import SqliteConnection from sqlspec.core.compiler import OperationType __all__ = ( + "SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", "SqliteStreamSource", "apply_driver_features", "build_connection_config", @@ -49,6 +51,7 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "end_transaction", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", @@ -75,8 +78,36 @@ SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT = sys.version_info >= (3, 12) SQLITE_DATABASE_LIST_MIN_COLUMNS = 2 SQLITE_TABLE_LIST_MIN_COLUMNS = 5 -SQLITE_TABLE_INFO_MIN_COLUMNS = 2 -SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") + + +def end_transaction( + connection: "SqliteConnection | Any", + *, + commit: bool, + supports_autocommit: bool = SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT, +) -> None: + """End an open transaction on a connection. + + Connection.commit and Connection.rollback are no-ops while the connection runs + in sqlite3's autocommit mode, so the statement is issued directly there. + + Args: + connection: Connection whose transaction should end. + commit: Whether to commit rather than roll back. + supports_autocommit: Whether this runtime's sqlite3 exposes autocommit. + """ + if not getattr(connection, "in_transaction", True): + return + if supports_autocommit and getattr(connection, "autocommit", None) is True: + connection.execute("COMMIT" if commit else "ROLLBACK") + return + if commit: + connection.commit() + else: + connection.rollback() + + +_end_transaction = end_transaction _TIME_TO_ISO = time_iso_convert @@ -221,13 +252,14 @@ def normalize_execute_parameters(parameters: Any) -> Any: class SqliteStreamSource: """Compiled chunk source streaming dict rows from a SQLite cursor via ``fetchmany``.""" - __slots__ = ("_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") + __slots__ = ("_as_dict", "_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") - def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: + def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int, as_dict: bool = True) -> None: self._driver = driver self._sql = sql self._parameters = parameters self._chunk_size = chunk_size + self._as_dict = as_dict self._cursor: Any = None self._column_names: list[str] | None = None @@ -240,7 +272,7 @@ def start(self) -> None: cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) self._driver._check_pending_exception(handler) - def fetch_chunk(self) -> "list[dict[str, Any]]": + def fetch_chunk(self) -> "list[Any]": handler = self._driver.handle_database_exceptions() rows: list[Any] = [] with handler: @@ -248,6 +280,8 @@ def fetch_chunk(self) -> "list[dict[str, Any]]": self._driver._check_pending_exception(handler) if not rows: return [] + if not self._as_dict: + return rows if self._column_names is None: self._column_names = [description[0] for description in self._cursor.description] return rows_to_dicts(rows, self._column_names) @@ -294,13 +328,17 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str "path", "file", } - connection_parameters = { - key: value - for key, value in connection_config.items() - if key not in excluded_keys - and (value is not None or key == "isolation_level") - and (key != "autocommit" or SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT) - } + + def _filter_params(mapping: Mapping[str, Any]) -> dict[str, Any]: + return { + key: value + for key, value in mapping.items() + if key not in excluded_keys + and (value is not None or key == "isolation_level") + and (key != "autocommit" or SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT) + } + + connection_parameters = _filter_params(connection_config) if "database" not in connection_parameters: database = ( @@ -314,13 +352,7 @@ def build_connection_config(connection_config: "Mapping[str, Any]") -> "dict[str extra = connection_config.get("extra") if isinstance(extra, Mapping): - connection_parameters.update({ - key: value - for key, value in extra.items() - if key not in excluded_keys - and (value is not None or key == "isolation_level") - and (key != "autocommit" or SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT) - }) + connection_parameters.update(_filter_params(extra)) return connection_parameters @@ -360,9 +392,6 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) ): return _create_sqlite_error(error, error_code, UniqueViolationError, "unique constraint violation") - # Check for busy/locked conditions first (deadlock-like scenarios in SQLite) - # SQLITE_BUSY means another process has the database locked - # SQLITE_LOCKED means another connection has the table/rows locked if error_code == SQLITE_BUSY_CODE or error_name == "SQLITE_BUSY": return _create_sqlite_error(error, error_code, DeadlockError, "database busy") if error_code == SQLITE_LOCKED_CODE or error_name == "SQLITE_LOCKED": @@ -370,13 +399,11 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) if "locked" in error_msg or "busy" in error_msg: return _create_sqlite_error(error, error_code or 0, DeadlockError, "database locked") - # Query interruption (timeout-like behavior) if error_code == SQLITE_INTERRUPT_CODE or error_name == "SQLITE_INTERRUPT": return _create_sqlite_error(error, error_code, OperationCancelledError, "query interrupted") if "interrupt" in error_msg: return _create_sqlite_error(error, error_code or 0, OperationCancelledError, "query interrupted") - # Permission errors if error_code == SQLITE_PERM_CODE or error_name == "SQLITE_PERM": return _create_sqlite_error(error, error_code, PermissionDeniedError, "permission denied") if error_code == SQLITE_READONLY_CODE or error_name == "SQLITE_READONLY": @@ -397,7 +424,6 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) return _create_sqlite_error(error, None, SQLParsingError, "SQL syntax error") return _create_sqlite_error(error, None, SQLSpecError, "database error") - # Constraint violations (check extended error codes first) if error_code == SQLITE_CONSTRAINT_FOREIGNKEY_CODE or error_name == "SQLITE_CONSTRAINT_FOREIGNKEY": return _create_sqlite_error(error, error_code, ForeignKeyViolationError, "foreign key constraint violation") if error_code == SQLITE_CONSTRAINT_NOTNULL_CODE or error_name == "SQLITE_CONSTRAINT_NOTNULL": @@ -407,17 +433,14 @@ def create_mapped_exception(error: BaseException, *, logger: Any | None = None) if error_code == SQLITE_CONSTRAINT_CODE or error_name == "SQLITE_CONSTRAINT": return _create_sqlite_error(error, error_code, IntegrityError, "integrity constraint violation") - # Connection/file errors if error_code == SQLITE_CANTOPEN_CODE or error_name == "SQLITE_CANTOPEN": return _create_sqlite_error(error, error_code, DatabaseConnectionError, "connection error") if error_code == SQLITE_IOERR_CODE or error_name == "SQLITE_IOERR": return _create_sqlite_error(error, error_code, OperationalError, "operational error") - # Data type errors if error_code == SQLITE_MISMATCH_CODE or error_name == "SQLITE_MISMATCH": return _create_sqlite_error(error, error_code, DataError, "data error") - # SQL syntax errors if error_code == 1 or "syntax" in error_msg: return _create_sqlite_error(error, error_code, SQLParsingError, "SQL syntax error") @@ -440,7 +463,7 @@ def build_profile() -> "DriverParameterProfile": preserve_original_params_for_many=False, json_serializer_strategy="helper", custom_type_coercions={ - bool: _bool_to_int, + bool: int, datetime: _TIME_TO_ISO, date: _TIME_TO_ISO, Decimal: _DECIMAL_TO_STRING, @@ -507,7 +530,7 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> with contextlib.suppress(Exception): table_cursor.close() except sqlite3.Error: - return _target_supports_rowid_legacy(connection, target) + return False candidates = [ row @@ -521,10 +544,10 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> if target_schema is not None: candidates = [row for row in candidates if row[0].casefold() == target_schema.casefold()] if not candidates: - return _target_supports_rowid_legacy(connection, target) + return False return len(candidates) == 1 and candidates[0][4] == 0 if not candidates: - return _target_supports_rowid_legacy(connection, target) + return False schema_order = ["temp", "main"] if not any(row[0].casefold() in {"temp", "main"} for row in candidates): @@ -549,80 +572,6 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> return False -def _target_supports_rowid_legacy(connection: Any, target: "tuple[str | None, str]") -> bool: - target_schema, target_table = target - schema_order = [target_schema] if target_schema is not None else ["temp", "main"] - if target_schema is None: - try: - database_cursor = connection.execute("PRAGMA database_list") - try: - schema_order.extend( - row[1] - for row in database_cursor.fetchall() - if len(row) >= SQLITE_DATABASE_LIST_MIN_COLUMNS - and isinstance(row[1], str) - and row[1] not in {"main", "temp"} - ) - finally: - with contextlib.suppress(Exception): - database_cursor.close() - except sqlite3.Error: - return False - - for schema_name in schema_order: - if schema_name is None: - continue - quoted_schema = quote_identifier(schema_name) - schema_cursor = None - schema_row = None - try: - schema_cursor = connection.execute( - f"SELECT type FROM {quoted_schema}.sqlite_master WHERE name = ? COLLATE NOCASE", (target_table,) - ) - schema_row = schema_cursor.fetchone() - except sqlite3.Error: - pass - finally: - if schema_cursor is not None: - with contextlib.suppress(Exception): - schema_cursor.close() - if schema_row is None: - continue - if not schema_row or schema_row[0] != "table": - return False - qualified_target = f"{quoted_schema}.{quote_identifier(target_table)}" - table_info_cursor = None - try: - table_info_cursor = connection.execute( - f"PRAGMA {quoted_schema}.table_info({quote_identifier(target_table)})" - ) - column_names = { - row[1].casefold() - for row in table_info_cursor.fetchall() - if len(row) >= SQLITE_TABLE_INFO_MIN_COLUMNS and isinstance(row[1], str) - } - except sqlite3.Error: - return False - finally: - if table_info_cursor is not None: - with contextlib.suppress(Exception): - table_info_cursor.close() - hidden_alias = next((alias for alias in SQLITE_ROWID_ALIASES if alias not in column_names), None) - if hidden_alias is None: - return False - probe_cursor = None - try: - probe_cursor = connection.execute(f"SELECT {hidden_alias} FROM {qualified_target} LIMIT 0") - except sqlite3.Error: - return False - finally: - if probe_cursor is not None: - with contextlib.suppress(Exception): - probe_cursor.close() - return True - return False - - def _create_sqlite_error( error: Any, code: "int | None", error_class: type[SQLSpecError], description: str ) -> SQLSpecError: @@ -644,10 +593,6 @@ def _create_sqlite_error( return exc -def _bool_to_int(value: bool) -> int: - return int(value) - - driver_profile = build_profile() default_statement_config = build_statement_config() diff --git a/sqlspec/adapters/sqlite/data_dictionary.py b/sqlspec/adapters/sqlite/data_dictionary.py index 7d1fa92f4..20ff4ac92 100644 --- a/sqlspec/adapters/sqlite/data_dictionary.py +++ b/sqlspec/adapters/sqlite/data_dictionary.py @@ -50,10 +50,8 @@ def get_version(self, driver: "SqliteDriver") -> "VersionInfo | None": SQLite version information or None if detection fails. """ driver_id = id(driver) - # Inline cache check to avoid cross-module method call that causes mypyc segfault if driver_id in self._version_fetch_attempted: return self._version_cache.get(driver_id) - # Not cached, fetch from database version_value = driver.select_value_or_none(self.get_query("version", "current")) if not version_value: diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index b8ec11f98..8b7844b7f 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -1,6 +1,6 @@ """SQLite driver implementation.""" -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast from sqlspec.adapters.sqlite._typing import SqliteCursor, SqliteSessionContext from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 @@ -44,6 +44,9 @@ __all__ = ("SqliteCursor", "SqliteDriver", "SqliteExceptionHandler", "SqliteSessionContext") +T = TypeVar("T") +_BATCH_SAMPLE_THRESHOLD: Final = 100 + class SqliteExceptionHandler(BaseSyncExceptionHandler): """Context manager for handling SQLite database exceptions. @@ -219,15 +222,20 @@ def execute_many( return DMLResult(operation, affected_rows) return super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - def begin(self) -> None: + def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: """Begin a database transaction. + Args: + mode: Transaction lock mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). + Defaults to configured driver feature or IMMEDIATE. + Raises: SQLSpecError: If transaction cannot be started """ + transaction_mode = mode or self.driver_features.get("default_transaction_mode", "IMMEDIATE") try: if not self.connection.in_transaction: - self.connection.execute("BEGIN") + self.connection.execute(f"BEGIN {transaction_mode}") except sqlite3.Error as e: msg = f"Failed to begin transaction: {e}" raise SQLSpecError(msg) from e @@ -284,12 +292,14 @@ def with_cursor(self, connection: "SqliteConnection") -> "SqliteCursor": """ return SqliteCursor(connection) - def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": + def dispatch_select_stream( + self, statement: "SQL", chunk_size: int, as_dict: bool = True + ) -> "SyncRowStream[Any] | None": """Return a native SQLite row stream backed by chunked ``fetchmany``.""" if not statement.returns_rows(): return None sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) - return SyncRowStream(SqliteStreamSource(self, sql, prepared_parameters, chunk_size)) + return SyncRowStream(SqliteStreamSource(self, sql, prepared_parameters, chunk_size, as_dict=as_dict)) def handle_database_exceptions(self) -> "SqliteExceptionHandler": """Handle database-specific exceptions and wrap them appropriately. @@ -312,7 +322,6 @@ def select_to_storage( **kwargs: Any, ) -> "StorageBridgeJob": """Execute a query and write Arrow-compatible output to storage (sync).""" - self._require_capability("arrow_export_enabled") arrow_result = self.select_to_arrow(statement, *parameters, statement_config=statement_config, **kwargs) sync_pipeline = self._storage_pipeline() @@ -327,37 +336,42 @@ def load_from_arrow( table: str, source: "ArrowResult | Any", *, + batch_size: int = 10000, partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, ) -> "StorageBridgeJob": - """Load Arrow data into SQLite using batched inserts.""" - + """Load Arrow data into SQLite using chunked 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 - ) + columns = arrow_table.column_names + insert_sql = build_insert_statement(table, columns) + needs_prep = self._arrow_rows_need_preparation(arrow_table) + owns_transaction = not self.connection.in_transaction try: if owns_transaction: - self.connection.execute("BEGIN IMMEDIATE") + self.begin("IMMEDIATE") if overwrite: statement = f"DELETE FROM {format_identifier(table)}" with self.with_cursor(self.connection) as cursor: cursor.execute(statement) - if records: - insert_sql = build_insert_statement(table, columns) - with self.with_cursor(self.connection) as cursor: - cursor.executemany(insert_sql, cast("Any", prepared_records)) + for batch in arrow_table.to_batches(max_chunksize=batch_size): + pydict = batch.to_pydict() + records = list(zip(*(pydict[col] for col in columns), strict=False)) + if records: + prepared_records = ( + self.prepare_driver_parameters(records, self.statement_config, is_many=True) + if needs_prep + else records + ) + with self.with_cursor(self.connection) as cursor: + cursor.executemany(insert_sql, cast("Any", prepared_records)) if owns_transaction: - self.connection.commit() + self.commit() except sqlite3.Error as exc: if owns_transaction: - self.connection.rollback() + self.rollback() raise create_mapped_exception(exc) from exc telemetry_payload = self._ingest_telemetry(arrow_table) @@ -430,14 +444,16 @@ def _execute_cache_hit( fetched_data = cursor.fetchall() affected_rows = resolve_rowcount(cursor) - last_inserted_id = resolve_lastrowid( - self.connection, - cursor, - cached.operation_type, - affected_rows, - cached.processed_state.parsed_expression, - self._rowid_target_cache, - ) + last_inserted_id = None + if cached.operation_type != "SELECT": + last_inserted_id = resolve_lastrowid( + self.connection, + cursor, + cached.operation_type, + affected_rows, + cached.processed_state.parsed_expression, + self._rowid_target_cache, + ) description = cursor.description column_names = [col[0] for col in description] if description else [] row_format = resolve_row_format(fetched_data) @@ -458,7 +474,7 @@ def _execute_cache_hit( if direct_statement is not None: self._release_pooled_statement(direct_statement) msg = "unreachable" - raise AssertionError(msg) # pragma: no cover + raise AssertionError(msg) def _invalidate_rowid_target_cache(self, operation_type: "OperationType") -> None: if operation_type not in {"SELECT", "INSERT", "UPDATE", "DELETE"}: @@ -506,10 +522,16 @@ def _thin_path_parameters_are_eligible( has_type_coercion = bool(coercion_map) fallback_items = type_coercion_fallbacks(coercion_map) if coercion_map else () - # Common benchmark shape: list[tuple[value]] + total_rows = len(parameters) + if total_rows > _BATCH_SAMPLE_THRESHOLD: + sample_indices = (0, 1, total_rows // 4, total_rows // 2, (3 * total_rows) // 4, total_rows - 1) + eval_parameters = [parameters[i] for i in sample_indices] + else: + eval_parameters = parameters + if row_len == 1: if has_type_coercion and coercion_map is not None: - for param_set in parameters: + for param_set in eval_parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -519,7 +541,7 @@ def _thin_path_parameters_are_eligible( return False return True - for param_set in parameters: + for param_set in eval_parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -530,7 +552,7 @@ def _thin_path_parameters_are_eligible( return True if has_type_coercion and coercion_map is not None: - for param_set in parameters: + for param_set in eval_parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -541,7 +563,7 @@ def _thin_path_parameters_are_eligible( return False return True - for param_set in parameters: + for param_set in eval_parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -571,13 +593,5 @@ def _resolve_dml_operation_type(statement: str) -> "OperationType": 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) - register_driver_profile("sqlite", driver_profile) diff --git a/sqlspec/adapters/sqlite/events/store.py b/sqlspec/adapters/sqlite/events/store.py index 1b5233d10..ef2da88e3 100644 --- a/sqlspec/adapters/sqlite/events/store.py +++ b/sqlspec/adapters/sqlite/events/store.py @@ -4,7 +4,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.sqlite.config import SqliteConfig, _apply_extension_pragmas, _extension_pragma_statements +from sqlspec.adapters.sqlite.config import SqliteConfig, apply_extension_pragmas, extension_pragma_statements from sqlspec.config import EventsConfig from sqlspec.extensions.events import BaseEventQueueStore @@ -36,11 +36,11 @@ class SqliteEventQueueStore(BaseEventQueueStore[SqliteConfig]): def __init__(self, config: SqliteConfig) -> None: super().__init__(config) - self._pragma_statements = _extension_pragma_statements(config, "events") + self._pragma_statements = extension_pragma_statements(config, "events") def prepare_schema_sync(self, driver: Any) -> None: """Apply configured SQLite PRAGMAs before queue DDL.""" - _apply_extension_pragmas(driver.connection, self._pragma_statements) + apply_extension_pragmas(driver.connection, self._pragma_statements) def _column_types(self) -> "tuple[str, str, str]": """Return SQLite-compatible column types for the event queue.""" diff --git a/sqlspec/adapters/sqlite/litestar/store.py b/sqlspec/adapters/sqlite/litestar/store.py index 8fbb0915b..b8875536d 100644 --- a/sqlspec/adapters/sqlite/litestar/store.py +++ b/sqlspec/adapters/sqlite/litestar/store.py @@ -5,7 +5,8 @@ from typing_extensions import NotRequired -from sqlspec.adapters.sqlite.config import _apply_extension_pragmas, _extension_pragma_statements +from sqlspec.adapters.sqlite.config import apply_extension_pragmas, extension_pragma_statements +from sqlspec.adapters.sqlite.core import end_transaction from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ @@ -63,7 +64,7 @@ def __init__(self, config: "SqliteConfig") -> None: config: SqliteConfig instance. """ super().__init__(config) - self._pragma_statements = _extension_pragma_statements(config, "litestar") + self._pragma_statements = extension_pragma_statements(config, "litestar") async def create_table(self) -> None: """Create the session table if it doesn't exist.""" @@ -75,7 +76,7 @@ async def create_table(self) -> None: def prepare_schema_sync(self, driver: Any) -> None: """Apply configured SQLite PRAGMAs before migration DDL generation.""" - _apply_extension_pragmas(driver.connection, self._pragma_statements) + apply_extension_pragmas(driver.connection, self._pragma_statements) async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Get a session value by key. @@ -201,7 +202,7 @@ def _create_table(self) -> None: """Synchronous implementation of create_table.""" sql = self._table_ddl() with self._config.provide_session() as driver: - _apply_extension_pragmas(driver.connection, self._pragma_statements) + apply_extension_pragmas(driver.connection, self._pragma_statements) driver.execute_script(sql) self._log_table_created() @@ -210,7 +211,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = ? - AND (expires_at IS NULL OR julianday(expires_at) > julianday('now')) + AND (expires_at IS NULL OR expires_at > julianday('now')) """ with self._config.provide_connection() as conn: @@ -232,7 +233,7 @@ def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | WHERE session_id = ? """ conn.execute(update_sql, (new_expires_at_julian, key)) - conn.commit() + end_transaction(conn, commit=True) return bytes(data) @@ -249,7 +250,7 @@ def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | No with self._config.provide_connection() as conn: conn.execute(sql, (key, data, expires_at_julian)) - conn.commit() + end_transaction(conn, commit=True) def _delete(self, key: str) -> None: """Synchronous implementation of delete.""" @@ -257,7 +258,7 @@ def _delete(self, key: str) -> None: with self._config.provide_connection() as conn: conn.execute(sql, (key,)) - conn.commit() + end_transaction(conn, commit=True) def _delete_all(self) -> None: """Synchronous implementation of delete_all.""" @@ -265,7 +266,7 @@ def _delete_all(self) -> None: with self._config.provide_connection() as conn: conn.execute(sql) - conn.commit() + end_transaction(conn, commit=True) self._log_delete_all() def _exists(self, key: str) -> bool: @@ -273,7 +274,7 @@ def _exists(self, key: str) -> bool: sql = f""" SELECT 1 FROM {self._table_name} WHERE session_id = ? - AND (expires_at IS NULL OR julianday(expires_at) > julianday('now')) + AND (expires_at IS NULL OR expires_at > julianday('now')) """ with self._config.provide_connection() as conn: @@ -311,11 +312,11 @@ def _expires_in(self, key: str) -> "int | None": def _delete_expired(self) -> int: """Synchronous implementation of delete_expired.""" - sql = f"DELETE FROM {self._table_name} WHERE julianday(expires_at) <= julianday('now')" + sql = f"DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= julianday('now')" with self._config.provide_connection() as conn: cursor = conn.execute(sql) - conn.commit() + end_transaction(conn, commit=True) count = cursor.rowcount if count > 0: self._log_delete_expired(count) diff --git a/sqlspec/adapters/sqlite/pool.py b/sqlspec/adapters/sqlite/pool.py index e18936e41..5c8595dc9 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -9,14 +9,21 @@ from sqlspec.adapters.sqlite._typing import SqliteConnection from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 -from sqlspec.adapters.sqlite.core import SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT +from sqlspec.adapters.sqlite.core import end_transaction from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 if TYPE_CHECKING: from collections.abc import Callable, Generator -__all__ = ("SqliteConnectionPool",) +__all__ = ( + "SQLITE_DISK_CACHE_SIZE", + "SQLITE_JOURNAL_SIZE_LIMIT", + "SQLITE_MMAP_SIZE", + "SqliteConnectionPool", + "_end_transaction", + "end_transaction", +) logger = get_logger(POOL_LOGGER_NAME) _ADAPTER_NAME = "sqlite" @@ -24,49 +31,33 @@ SQLITE_DEFAULT_ENABLE_FOREIGN_KEYS: Final = False SQLITE_DEFAULT_ENABLE_OPTIMIZATIONS: Final = True SQLITE_MEMORY_CACHE_SIZE: Final = -16000 +SQLITE_DISK_CACHE_SIZE: Final = -64000 +SQLITE_MMAP_SIZE: Final = 268435456 +SQLITE_JOURNAL_SIZE_LIMIT: Final = 67108864 SQLITE_WAL_SWITCH_ATTEMPTS: Final = 50 SQLITE_WAL_SWITCH_DELAY: Final = 0.01 +def _attempt_wal_switch(connection: "SqliteConnection", attempt: int) -> bool: + """Attempt a single WAL mode switch, returning True on success.""" + try: + connection.execute("PRAGMA journal_mode = WAL") + except sqlite3.OperationalError as exc: + if "locked" not in str(exc) or attempt == SQLITE_WAL_SWITCH_ATTEMPTS - 1: + raise + time.sleep(SQLITE_WAL_SWITCH_DELAY) + return False + return True + + def _enable_wal(connection: "SqliteConnection") -> None: """Retry database and table locks briefly while switching to WAL mode.""" for attempt in range(SQLITE_WAL_SWITCH_ATTEMPTS): - try: - connection.execute("PRAGMA journal_mode = WAL") - except sqlite3.OperationalError as exc: # noqa: PERF203 - bounded lock retry - if "locked" not in str(exc) or attempt == SQLITE_WAL_SWITCH_ATTEMPTS - 1: - raise - time.sleep(SQLITE_WAL_SWITCH_DELAY) - else: + if _attempt_wal_switch(connection, attempt): return -def _end_transaction( - connection: SqliteConnection, *, commit: bool, supports_autocommit: bool = SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT -) -> None: - """End an open transaction on a connection. - - ``Connection.commit`` and ``Connection.rollback`` are no-ops while the - connection runs in sqlite3's autocommit mode, so the statement is issued - directly there. - - Args: - connection: Connection whose transaction should end. - commit: Whether to commit rather than roll back. - supports_autocommit: Whether this runtime's sqlite3 exposes autocommit. - - Returns: - None. - """ - if not connection.in_transaction: - return - if supports_autocommit and cast("Any", connection).autocommit is True: - connection.execute("COMMIT" if commit else "ROLLBACK") - return - if commit: - connection.commit() - else: - connection.rollback() +_end_transaction = end_transaction class SqliteConnectionPool: @@ -84,6 +75,7 @@ class SqliteConnectionPool: "_enable_optimizations", "_generation", "_health_check_interval", + "_is_memory_db", "_on_connection_create", "_pool_id", "_recycle_seconds", @@ -116,6 +108,8 @@ def __init__( if "check_same_thread" not in connection_parameters: connection_parameters = {**connection_parameters, "check_same_thread": False} self._connection_parameters = connection_parameters + database = self._connection_parameters.get("database", ":memory:") + self._is_memory_db = database == ":memory:" or "mode=memory" in str(database) self._thread_local = threading.local() self._connection_registry: set[SqliteConnection] = set() self._generation = 0 @@ -155,20 +149,28 @@ def new_connection(self) -> SqliteConnection: connection = sqlite3.connect(**self._connection_parameters) try: - if self._enable_optimizations: - database = self._connection_parameters.get("database", ":memory:") - is_memory = database == ":memory:" or "mode=memory" in str(database) + busy_timeout = self._connection_parameters.get("busy_timeout", SQLITE_BUSY_TIMEOUT) + connection.execute(f"PRAGMA busy_timeout = {busy_timeout}") - if is_memory: + if self._enable_optimizations: + if self._is_memory_db: connection.execute("PRAGMA journal_mode = MEMORY") connection.execute("PRAGMA synchronous = OFF") connection.execute("PRAGMA temp_store = MEMORY") - connection.execute(f"PRAGMA cache_size = {SQLITE_MEMORY_CACHE_SIZE}") + cache_size = self._connection_parameters.get("cache_size", SQLITE_MEMORY_CACHE_SIZE) + connection.execute(f"PRAGMA cache_size = {cache_size}") else: - _enable_wal(connection) + current_mode = connection.execute("PRAGMA journal_mode").fetchone() + current_mode_str = str(current_mode[0]).lower() if current_mode else "" + if current_mode_str != "wal": + _enable_wal(connection) connection.execute("PRAGMA synchronous = NORMAL") - - connection.execute(f"PRAGMA busy_timeout = {SQLITE_BUSY_TIMEOUT}") + cache_size = self._connection_parameters.get("cache_size", SQLITE_DISK_CACHE_SIZE) + connection.execute(f"PRAGMA cache_size = {cache_size}") + mmap_size = self._connection_parameters.get("mmap_size", SQLITE_MMAP_SIZE) + connection.execute(f"PRAGMA mmap_size = {mmap_size}") + connection.execute("PRAGMA temp_store = MEMORY") + connection.execute(f"PRAGMA journal_size_limit = {SQLITE_JOURNAL_SIZE_LIMIT}") if self._enable_foreign_keys: connection.execute("PRAGMA foreign_keys = ON") @@ -183,7 +185,7 @@ def new_connection(self) -> SqliteConnection: connection.close() raise - return connection # type: ignore[no-any-return] + return cast("SqliteConnection", connection) def _is_connection_alive(self, connection: SqliteConnection) -> bool: """Check if a connection is still alive and usable. @@ -202,21 +204,27 @@ def _is_connection_alive(self, connection: SqliteConnection) -> bool: def _get_thread_connection(self) -> SqliteConnection: """Get or create a connection for the current thread.""" - thread_state = self._thread_local.__dict__ - if thread_state.get("generation") != self._generation: - stale = thread_state.pop("connection", None) + current_generation = getattr(self._thread_local, "generation", None) + if current_generation != self._generation: + stale = getattr(self._thread_local, "connection", None) if stale is not None: self._retire_connection(cast("SqliteConnection", stale)) - thread_state.pop("created_at", None) - thread_state.pop("last_used", None) + self._thread_local.connection = None + self._thread_local.created_at = 0.0 + self._thread_local.last_used = 0.0 self._thread_local.generation = self._generation - if "connection" not in thread_state: - self._thread_local.connection = self._create_connection() - self._thread_local.created_at = time.time() - self._thread_local.last_used = time.time() - return cast("SqliteConnection", self._thread_local.connection) - if self._recycle_seconds > 0 and time.time() - self._thread_local.created_at > self._recycle_seconds: + conn = getattr(self._thread_local, "connection", None) + now = time.time() + if conn is None: + conn = self._create_connection() + self._thread_local.connection = conn + self._thread_local.created_at = now + self._thread_local.last_used = now + return conn + + created_at = getattr(self._thread_local, "created_at", 0.0) + if self._recycle_seconds > 0 and (now - created_at) > self._recycle_seconds: log_with_context( logger, logging.DEBUG, @@ -227,14 +235,20 @@ def _get_thread_connection(self) -> SqliteConnection: recycle_seconds=self._recycle_seconds, reason="exceeded_recycle_time", ) - self._retire_connection(self._thread_local.connection) - self._thread_local.connection = self._create_connection() - self._thread_local.created_at = time.time() - self._thread_local.last_used = time.time() - return cast("SqliteConnection", self._thread_local.connection) - - idle_time = time.time() - thread_state.get("last_used", 0) - if idle_time > self._health_check_interval and not self._is_connection_alive(self._thread_local.connection): + self._retire_connection(cast("SqliteConnection", conn)) + conn = self._create_connection() + self._thread_local.connection = conn + self._thread_local.created_at = now + self._thread_local.last_used = now + return conn + + last_used = getattr(self._thread_local, "last_used", 0.0) + idle_time = now - last_used + if ( + not self._is_memory_db + and idle_time > self._health_check_interval + and not self._is_connection_alive(cast("SqliteConnection", conn)) + ): log_with_context( logger, logging.DEBUG, @@ -245,12 +259,15 @@ def _get_thread_connection(self) -> SqliteConnection: idle_seconds=round(idle_time, 1), reason="failed_health_check", ) - self._retire_connection(self._thread_local.connection) - self._thread_local.connection = self._create_connection() - self._thread_local.created_at = time.time() + self._retire_connection(cast("SqliteConnection", conn)) + conn = self._create_connection() + self._thread_local.connection = conn + self._thread_local.created_at = now + self._thread_local.last_used = now + return conn - self._thread_local.last_used = time.time() - return cast("SqliteConnection", self._thread_local.connection) + self._thread_local.last_used = now + return cast("SqliteConnection", conn) def _retire_connection(self, connection: SqliteConnection) -> None: """Close a pool-owned connection and drop it from the shutdown registry.""" @@ -261,14 +278,12 @@ def _retire_connection(self, connection: SqliteConnection) -> None: def _close_thread_connection(self) -> None: """Close the connection for the current thread.""" - thread_state = self._thread_local.__dict__ - if "connection" in thread_state: - self._retire_connection(cast("SqliteConnection", self._thread_local.connection)) - del self._thread_local.connection - if "created_at" in thread_state: - del self._thread_local.created_at - if "last_used" in thread_state: - del self._thread_local.last_used + conn = getattr(self._thread_local, "connection", None) + if conn is not None: + self._retire_connection(cast("SqliteConnection", conn)) + self._thread_local.connection = None + self._thread_local.created_at = 0.0 + self._thread_local.last_used = 0.0 @contextmanager def get_connection(self) -> "Generator[SqliteConnection, None, None]": @@ -316,12 +331,9 @@ def release(self, connection: SqliteConnection) -> None: def size(self) -> int: """Get pool size (always 1 for thread-local).""" - try: - _ = self._thread_local.connection - except AttributeError: - return 0 - else: + if getattr(self._thread_local, "connection", None) is not None: return 1 + return 0 def checked_out(self) -> int: """Get number of checked out connections (always 0).""" @@ -366,7 +378,7 @@ def _apply_runtime_setup(connection: SqliteConnection, runtime_setup: "dict[str, function_config["name"], function_config["narg"], function_config["func"], - deterministic=function_config.get("deterministic", False), + deterministic=function_config.get("deterministic", True), ) for aggregate_config in runtime_setup.get("custom_aggregates", ()): @@ -374,6 +386,11 @@ def _apply_runtime_setup(connection: SqliteConnection, runtime_setup: "dict[str, aggregate_config["name"], aggregate_config["narg"], aggregate_config["aggregate_class"] ) + create_window_fn = getattr(connection, "create_window_function", None) + if create_window_fn is not None: + for window_config in runtime_setup.get("custom_window_functions", ()): + create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) + for collation_config in runtime_setup.get("custom_collations", ()): connection.create_collation(collation_config["name"], collation_config["func"]) diff --git a/sqlspec/adapters/sqlite/type_converter.py b/sqlspec/adapters/sqlite/type_converter.py index fad74a2e3..c113759e9 100644 --- a/sqlspec/adapters/sqlite/type_converter.py +++ b/sqlspec/adapters/sqlite/type_converter.py @@ -1,4 +1,3 @@ -# Keep in sync with sqlspec/adapters/aiosqlite/type_converter.py """SQLite custom type handlers for optional JSON and type conversion support. Provides registration functions for SQLite's adapter/converter system to enable @@ -9,12 +8,12 @@ instead of lambdas for adapter registration. """ -import json from functools import partial from typing import TYPE_CHECKING, Any from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 from sqlspec.utils.logging import get_logger +from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Callable @@ -31,14 +30,13 @@ def json_adapter(value: Any, serializer: "Callable[[Any], str] | None" = None) - Args: value: Python dict or list to serialize. - serializer: Optional JSON serializer callable. Defaults to standard json.dumps. + serializer: Optional JSON serializer callable. Defaults to compiled to_json. Returns: JSON string representation. """ - if serializer is None: - return json.dumps(value, ensure_ascii=False) - return serializer(value) + codec = serializer or to_json + return codec(value) def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = None) -> Any: @@ -46,14 +44,13 @@ def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = N Args: value: UTF-8 encoded JSON bytes from SQLite. - deserializer: Optional JSON deserializer callable. Defaults to standard json.loads. + deserializer: Optional JSON deserializer callable. Defaults to compiled from_json. Returns: Deserialized Python object (dict or list). """ - if deserializer is None: - return json.loads(value.decode("utf-8")) - return deserializer(value.decode("utf-8")) + codec = deserializer or from_json + return codec(value.decode("utf-8")) def register_type_handlers( @@ -64,6 +61,9 @@ def register_type_handlers( This function registers handlers globally for the sqlite3 module. It should be called once during application initialization if custom type handling is needed. + Note that sqlite3.register_adapter is deprecated in Python 3.12+ in favor of + statement-level conversions. + Args: json_serializer: Optional custom JSON serializer. json_deserializer: Optional custom JSON deserializer. diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 66bb8a7ca..16b5b291d 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -257,6 +257,8 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", + "custom_window_functions", + "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", From 3512cfb4c6d7777c2be979754a24e53c4d79213e Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 01:47:49 +0000 Subject: [PATCH 02/12] fix(adapters): restore sqlite and aiosqlite runtime compatibility and test contracts --- sqlspec/adapters/aiosqlite/core.py | 82 +++++++++++++++++++++++++++- sqlspec/adapters/aiosqlite/driver.py | 21 +++++-- sqlspec/adapters/aiosqlite/pool.py | 21 +------ sqlspec/adapters/sqlite/adk/store.py | 26 +++++++-- sqlspec/adapters/sqlite/config.py | 16 ------ sqlspec/adapters/sqlite/core.py | 82 +++++++++++++++++++++++++++- sqlspec/adapters/sqlite/driver.py | 15 ++++- sqlspec/adapters/sqlite/pool.py | 38 +++---------- 8 files changed, 215 insertions(+), 86 deletions(-) diff --git a/sqlspec/adapters/aiosqlite/core.py b/sqlspec/adapters/aiosqlite/core.py index af1af3297..94465842f 100644 --- a/sqlspec/adapters/aiosqlite/core.py +++ b/sqlspec/adapters/aiosqlite/core.py @@ -87,6 +87,8 @@ SQLITE_READONLY_CODE = 8 SQLITE_DATABASE_LIST_MIN_COLUMNS = 2 SQLITE_TABLE_LIST_MIN_COLUMNS = 5 +SQLITE_TABLE_INFO_MIN_COLUMNS = 2 +SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") async def run_on_worker_thread( @@ -590,7 +592,7 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> with contextlib.suppress(Exception): table_cursor.close() except sqlite3.Error: - return False + return _target_supports_rowid_legacy(connection, target) candidates = [ row @@ -604,10 +606,10 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> if target_schema is not None: candidates = [row for row in candidates if row[0].casefold() == target_schema.casefold()] if not candidates: - return False + return _target_supports_rowid_legacy(connection, target) return len(candidates) == 1 and candidates[0][4] == 0 if not candidates: - return False + return _target_supports_rowid_legacy(connection, target) schema_order = ["temp", "main"] if not any(row[0].casefold() in {"temp", "main"} for row in candidates): @@ -632,6 +634,80 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> return False +def _target_supports_rowid_legacy(connection: Any, target: "tuple[str | None, str]") -> bool: + target_schema, target_table = target + schema_order = [target_schema] if target_schema is not None else ["temp", "main"] + if target_schema is None: + try: + database_cursor = connection.execute("PRAGMA database_list") + try: + schema_order.extend( + row[1] + for row in database_cursor.fetchall() + if len(row) >= SQLITE_DATABASE_LIST_MIN_COLUMNS + and isinstance(row[1], str) + and row[1] not in {"main", "temp"} + ) + finally: + with contextlib.suppress(Exception): + database_cursor.close() + except sqlite3.Error: + return False + + for schema_name in schema_order: + if schema_name is None: + continue + quoted_schema = quote_identifier(schema_name) + schema_cursor = None + schema_row = None + try: + schema_cursor = connection.execute( + f"SELECT type FROM {quoted_schema}.sqlite_master WHERE name = ? COLLATE NOCASE", (target_table,) + ) + schema_row = schema_cursor.fetchone() + except sqlite3.Error: + pass + finally: + if schema_cursor is not None: + with contextlib.suppress(Exception): + schema_cursor.close() + if schema_row is None: + continue + if not schema_row or schema_row[0] != "table": + return False + qualified_target = f"{quoted_schema}.{quote_identifier(target_table)}" + table_info_cursor = None + try: + table_info_cursor = connection.execute( + f"PRAGMA {quoted_schema}.table_info({quote_identifier(target_table)})" + ) + column_names = { + row[1].casefold() + for row in table_info_cursor.fetchall() + if len(row) >= SQLITE_TABLE_INFO_MIN_COLUMNS and isinstance(row[1], str) + } + except sqlite3.Error: + return False + finally: + if table_info_cursor is not None: + with contextlib.suppress(Exception): + table_info_cursor.close() + hidden_alias = next((alias for alias in SQLITE_ROWID_ALIASES if alias not in column_names), None) + if hidden_alias is None: + return False + probe_cursor = None + try: + probe_cursor = connection.execute(f"SELECT {hidden_alias} FROM {qualified_target} LIMIT 0") + except sqlite3.Error: + return False + finally: + if probe_cursor is not None: + with contextlib.suppress(Exception): + probe_cursor.close() + return True + return False + + def _create_aiosqlite_error( error: Any, code: "int | None", error_class: type[SQLSpecError], description: str ) -> SQLSpecError: diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index 0cf175f56..cfca451ba 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -1,6 +1,8 @@ """AIOSQLite driver implementation for async SQLite operations.""" import asyncio +import contextlib +import inspect import secrets from typing import TYPE_CHECKING, Any, cast @@ -55,6 +57,10 @@ "AiosqliteSessionContext", ) +random = secrets.SystemRandom() +_execute_and_resolve_metadata = execute_and_resolve_metadata +_execute_fetchall_with_metadata = execute_fetchall_with_metadata + class AiosqliteExceptionHandler(BaseAsyncExceptionHandler): """Async context manager for handling aiosqlite database exceptions. @@ -107,7 +113,7 @@ async def dispatch_execute(self, cursor: "AiosqliteRawCursor", statement: "SQL") if statement.returns_rows(): fetched_data, description, _affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - execute_fetchall_with_metadata, + _execute_fetchall_with_metadata, self.connection, sql, normalized_parameters, @@ -131,7 +137,7 @@ async def dispatch_execute(self, cursor: "AiosqliteRawCursor", statement: "SQL") affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - execute_and_resolve_metadata, + _execute_and_resolve_metadata, self.connection, sql, normalized_parameters, @@ -204,7 +210,10 @@ async def execute_many( raise create_mapped_exception(exc) from exc finally: if cursor is not None: - await cursor.close() + with contextlib.suppress(Exception): + close_result = cursor.close() + if inspect.isawaitable(close_result): + await close_result return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) async def begin(self) -> None: @@ -369,7 +378,7 @@ async def _execute_cache_hit( 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, + _execute_fetchall_with_metadata, self.connection, cached.compiled_sql, normalized_parameters, @@ -399,7 +408,7 @@ async def _execute_cache_hit( else: affected_rows, last_inserted_id = await run_on_worker_thread( self.connection, - execute_and_resolve_metadata, + _execute_and_resolve_metadata, self.connection, cached.compiled_sql, normalized_parameters, @@ -553,7 +562,7 @@ async def _retry_begin_with_backoff( SQLSpecError: If every retry attempt fails. """ for attempt in range(max_retries): - delay = 0.01 * (2**attempt) + secrets.SystemRandom().uniform(0, 0.01) + delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) await asyncio.sleep(delay) try: await connection.execute("BEGIN IMMEDIATE") diff --git a/sqlspec/adapters/aiosqlite/pool.py b/sqlspec/adapters/aiosqlite/pool.py index e65594b04..23e577bb9 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -22,9 +22,6 @@ from sqlspec.adapters.aiosqlite._typing import AiosqliteConnection __all__ = ( - "SQLITE_DISK_CACHE_SIZE", - "SQLITE_JOURNAL_SIZE_LIMIT", - "SQLITE_MMAP_SIZE", "AiosqliteConnectTimeoutError", "AiosqliteConnectionPool", "AiosqlitePoolClosedError", @@ -39,9 +36,6 @@ SQLITE_DEFAULT_ENABLE_FOREIGN_KEYS: Final = False SQLITE_DEFAULT_ENABLE_OPTIMIZATIONS: Final = True SQLITE_MEMORY_CACHE_SIZE: Final = -16000 -SQLITE_DISK_CACHE_SIZE: Final = -64000 -SQLITE_MMAP_SIZE: Final = 268435456 -SQLITE_JOURNAL_SIZE_LIMIT: Final = 67108864 SQLITE_WAL_SWITCH_ATTEMPTS: Final = 50 SQLITE_WAL_SWITCH_DELAY: Final = 0.01 @@ -519,19 +513,8 @@ async def new_connection(self) -> "AiosqliteConnection": f"PRAGMA cache_size = {SQLITE_MEMORY_CACHE_SIZE}", ]) else: - cursor = await connection.execute("PRAGMA journal_mode") - current_mode = await cursor.fetchone() - await cursor.close() - if not current_mode or str(current_mode[0]).upper() != "WAL": - await _enable_wal(connection) - pragma_lines.extend([ - "PRAGMA synchronous = NORMAL", - "PRAGMA temp_store = MEMORY", - f"PRAGMA mmap_size = {SQLITE_MMAP_SIZE}", - f"PRAGMA cache_size = {SQLITE_DISK_CACHE_SIZE}", - f"PRAGMA journal_size_limit = {SQLITE_JOURNAL_SIZE_LIMIT}", - "PRAGMA threads = 4", - ]) + await _enable_wal(connection) + pragma_lines.append("PRAGMA synchronous = NORMAL") pragma_lines.append(f"PRAGMA busy_timeout = {SQLITE_BUSY_TIMEOUT}") diff --git a/sqlspec/adapters/sqlite/adk/store.py b/sqlspec/adapters/sqlite/adk/store.py index debb39e2e..1f39a1e3a 100644 --- a/sqlspec/adapters/sqlite/adk/store.py +++ b/sqlspec/adapters/sqlite/adk/store.py @@ -85,7 +85,6 @@ def __init__(self, config: "SqliteConfig") -> None: def create_tables(self) -> None: """Create both sessions and events tables if they don't exist.""" - """Synchronous implementation of create_tables.""" if not self.create_schema_enabled: self.reconcile_schema() return @@ -113,7 +112,6 @@ def create_session( Returns: Created session record. """ - """Synchronous implementation of create_session.""" now = datetime.now(timezone.utc) now_julian = _datetime_to_julian(now) state_json = to_json(state) @@ -134,6 +132,7 @@ def create_session( params = (session_id, app_name, user_id, state_json, now_julian, now_julian) with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, params) end_transaction(conn, commit=True) @@ -177,6 +176,7 @@ def get_session( try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) if update_sql: conn.execute(update_sql, update_params) end_transaction(conn, commit=True) @@ -218,6 +218,7 @@ def update_session_state(self, app_name: str, user_id: str, session_id: str, sta """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, (state_json, now_julian, app_name, user_id, session_id)) end_transaction(conn, commit=True) @@ -254,6 +255,7 @@ def list_sessions( try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, params) rows = cursor.fetchall() @@ -284,6 +286,7 @@ def delete_session(self, app_name: str, user_id: str, session_id: str) -> None: sql = f"DELETE FROM {self._session_table} WHERE app_name = ? AND user_id = ? AND id = ?" with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, (app_name, user_id, session_id)) end_transaction(conn, commit=True) @@ -303,6 +306,7 @@ def append_event(self, event_record: StoredEvent) -> None: """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute( sql, ( @@ -378,6 +382,7 @@ def append_event_and_update_state( """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) try: cursor = conn.execute(update_sql, (state_json, now_julian, app_name, user_id, session_id)) row = cursor.fetchone() @@ -440,7 +445,6 @@ def get_events( Returns: List of event records ordered by timestamp ASC. """ - """Synchronous implementation of get_events.""" if limit == 0: return [] @@ -463,6 +467,7 @@ def get_events( try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, params) rows = cursor.fetchall() @@ -493,6 +498,7 @@ def delete_expired_events(self, before: datetime, app_name: "str | None" = None) try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 end_transaction(conn, commit=True) @@ -512,6 +518,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 end_transaction(conn, commit=True) @@ -531,6 +538,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 end_transaction(conn, commit=True) @@ -546,6 +554,7 @@ def get_app_state(self, app_name: str) -> "dict[str, Any] | None": try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, (app_name,)) row = cursor.fetchone() return from_json(row[0]) if row is not None and row[0] else None @@ -564,6 +573,7 @@ def get_user_state(self, app_name: str, user_id: str) -> "dict[str, Any] | None" try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, (app_name, user_id)) row = cursor.fetchone() return from_json(row[0]) if row is not None and row[0] else None @@ -583,6 +593,7 @@ def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None: """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, (app_name, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))) end_transaction(conn, commit=True) @@ -597,6 +608,7 @@ def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]" """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, (app_name, user_id, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))) end_transaction(conn, commit=True) @@ -606,6 +618,7 @@ def get_metadata(self, key: str) -> "str | None": try: with self._config.provide_connection() as conn: + self._apply_pragmas(conn) cursor = conn.execute(sql, (key,)) row = cursor.fetchone() return str(row[0]) if row is not None else None @@ -623,6 +636,7 @@ def set_metadata(self, key: str, value: str) -> None: """ with self._config.provide_connection() as conn: + self._apply_pragmas(conn) conn.execute(sql, (key, value)) end_transaction(conn, commit=True) @@ -632,6 +646,10 @@ def _apply_pragmas(self, connection: Any) -> None: Args: connection: SQLite connection. """ + connection.execute("PRAGMA foreign_keys = ON") + connection.execute("PRAGMA cache_size = -64000") + connection.execute("PRAGMA mmap_size = 30000000") + connection.execute("PRAGMA journal_size_limit = 67108864") for pragma_name, pragma_value in self._pragma_overrides: connection.execute(f"PRAGMA {pragma_name} = {pragma_value}") @@ -765,7 +783,6 @@ def __init__(self, config: "SqliteConfig") -> None: self._fts_options = _fts_options(config) def create_tables(self) -> None: - """Create tables if they don't exist.""" """Create the memory table and indexes if they don't exist. Skips table creation if memory store is disabled. @@ -782,7 +799,6 @@ def create_tables(self) -> None: driver.execute_script(self._memory_table_ddl()) def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: "object | None" = None) -> int: - """Bulk insert memory entries with deduplication.""" """Bulk insert memory entries with deduplication. Uses INSERT OR IGNORE to skip duplicates based on event_id diff --git a/sqlspec/adapters/sqlite/config.py b/sqlspec/adapters/sqlite/config.py index b529653d7..b31efff57 100644 --- a/sqlspec/adapters/sqlite/config.py +++ b/sqlspec/adapters/sqlite/config.py @@ -62,10 +62,6 @@ class SqliteConnectionParams(TypedDict): health_check_interval: NotRequired[float] enable_optimizations: NotRequired[bool] enable_foreign_keys: NotRequired[bool] - busy_timeout: NotRequired[int] - cache_size: NotRequired[int] - mmap_size: NotRequired[int] - default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] extra: NotRequired[dict[str, Any]] @@ -351,18 +347,6 @@ def _create_pool(self) -> SqliteConnectionPool: if enable_foreign_keys is not None: pool_kwargs["enable_foreign_keys"] = enable_foreign_keys - busy_timeout = self.connection_config.get("busy_timeout") - if busy_timeout is not None: - pool_kwargs["busy_timeout"] = busy_timeout - - cache_size = self.connection_config.get("cache_size") - if cache_size is not None: - pool_kwargs["cache_size"] = cache_size - - mmap_size = self.connection_config.get("mmap_size") - if mmap_size is not None: - pool_kwargs["mmap_size"] = mmap_size - pool = SqliteConnectionPool( connection_parameters=config_dict, on_connection_create=self._user_connection_hook, diff --git a/sqlspec/adapters/sqlite/core.py b/sqlspec/adapters/sqlite/core.py index af4a02a4e..dacbc498a 100644 --- a/sqlspec/adapters/sqlite/core.py +++ b/sqlspec/adapters/sqlite/core.py @@ -78,6 +78,8 @@ SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT = sys.version_info >= (3, 12) SQLITE_DATABASE_LIST_MIN_COLUMNS = 2 SQLITE_TABLE_LIST_MIN_COLUMNS = 5 +SQLITE_TABLE_INFO_MIN_COLUMNS = 2 +SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") def end_transaction( @@ -530,7 +532,7 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> with contextlib.suppress(Exception): table_cursor.close() except sqlite3.Error: - return False + return _target_supports_rowid_legacy(connection, target) candidates = [ row @@ -544,10 +546,10 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> if target_schema is not None: candidates = [row for row in candidates if row[0].casefold() == target_schema.casefold()] if not candidates: - return False + return _target_supports_rowid_legacy(connection, target) return len(candidates) == 1 and candidates[0][4] == 0 if not candidates: - return False + return _target_supports_rowid_legacy(connection, target) schema_order = ["temp", "main"] if not any(row[0].casefold() in {"temp", "main"} for row in candidates): @@ -572,6 +574,80 @@ def _target_supports_rowid(connection: Any, target: "tuple[str | None, str]") -> return False +def _target_supports_rowid_legacy(connection: Any, target: "tuple[str | None, str]") -> bool: + target_schema, target_table = target + schema_order = [target_schema] if target_schema is not None else ["temp", "main"] + if target_schema is None: + try: + database_cursor = connection.execute("PRAGMA database_list") + try: + schema_order.extend( + row[1] + for row in database_cursor.fetchall() + if len(row) >= SQLITE_DATABASE_LIST_MIN_COLUMNS + and isinstance(row[1], str) + and row[1] not in {"main", "temp"} + ) + finally: + with contextlib.suppress(Exception): + database_cursor.close() + except sqlite3.Error: + return False + + for schema_name in schema_order: + if schema_name is None: + continue + quoted_schema = quote_identifier(schema_name) + schema_cursor = None + schema_row = None + try: + schema_cursor = connection.execute( + f"SELECT type FROM {quoted_schema}.sqlite_master WHERE name = ? COLLATE NOCASE", (target_table,) + ) + schema_row = schema_cursor.fetchone() + except sqlite3.Error: + pass + finally: + if schema_cursor is not None: + with contextlib.suppress(Exception): + schema_cursor.close() + if schema_row is None: + continue + if not schema_row or schema_row[0] != "table": + return False + qualified_target = f"{quoted_schema}.{quote_identifier(target_table)}" + table_info_cursor = None + try: + table_info_cursor = connection.execute( + f"PRAGMA {quoted_schema}.table_info({quote_identifier(target_table)})" + ) + column_names = { + row[1].casefold() + for row in table_info_cursor.fetchall() + if len(row) >= SQLITE_TABLE_INFO_MIN_COLUMNS and isinstance(row[1], str) + } + except sqlite3.Error: + return False + finally: + if table_info_cursor is not None: + with contextlib.suppress(Exception): + table_info_cursor.close() + hidden_alias = next((alias for alias in SQLITE_ROWID_ALIASES if alias not in column_names), None) + if hidden_alias is None: + return False + probe_cursor = None + try: + probe_cursor = connection.execute(f"SELECT {hidden_alias} FROM {qualified_target} LIMIT 0") + except sqlite3.Error: + return False + finally: + if probe_cursor is not None: + with contextlib.suppress(Exception): + probe_cursor.close() + return True + return False + + def _create_sqlite_error( error: Any, code: "int | None", error_class: type[SQLSpecError], description: str ) -> SQLSpecError: diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index 8b7844b7f..3141e41f6 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -227,15 +227,16 @@ def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = N Args: mode: Transaction lock mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). - Defaults to configured driver feature or IMMEDIATE. + Defaults to configured driver feature or SQLite default (DEFERRED). Raises: SQLSpecError: If transaction cannot be started """ - transaction_mode = mode or self.driver_features.get("default_transaction_mode", "IMMEDIATE") + transaction_mode = mode or self.driver_features.get("default_transaction_mode") try: if not self.connection.in_transaction: - self.connection.execute(f"BEGIN {transaction_mode}") + stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" + self.connection.execute(stmt) except sqlite3.Error as e: msg = f"Failed to begin transaction: {e}" raise SQLSpecError(msg) from e @@ -593,5 +594,13 @@ def _resolve_dml_operation_type(statement: str) -> "OperationType": 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) + register_driver_profile("sqlite", driver_profile) diff --git a/sqlspec/adapters/sqlite/pool.py b/sqlspec/adapters/sqlite/pool.py index 5c8595dc9..aae34d70d 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -9,21 +9,14 @@ from sqlspec.adapters.sqlite._typing import SqliteConnection from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 -from sqlspec.adapters.sqlite.core import end_transaction +from sqlspec.adapters.sqlite.core import end_transaction as _end_transaction from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 if TYPE_CHECKING: from collections.abc import Callable, Generator -__all__ = ( - "SQLITE_DISK_CACHE_SIZE", - "SQLITE_JOURNAL_SIZE_LIMIT", - "SQLITE_MMAP_SIZE", - "SqliteConnectionPool", - "_end_transaction", - "end_transaction", -) +__all__ = ("SqliteConnectionPool",) logger = get_logger(POOL_LOGGER_NAME) _ADAPTER_NAME = "sqlite" @@ -31,9 +24,6 @@ SQLITE_DEFAULT_ENABLE_FOREIGN_KEYS: Final = False SQLITE_DEFAULT_ENABLE_OPTIMIZATIONS: Final = True SQLITE_MEMORY_CACHE_SIZE: Final = -16000 -SQLITE_DISK_CACHE_SIZE: Final = -64000 -SQLITE_MMAP_SIZE: Final = 268435456 -SQLITE_JOURNAL_SIZE_LIMIT: Final = 67108864 SQLITE_WAL_SWITCH_ATTEMPTS: Final = 50 SQLITE_WAL_SWITCH_DELAY: Final = 0.01 @@ -57,9 +47,6 @@ def _enable_wal(connection: "SqliteConnection") -> None: return -_end_transaction = end_transaction - - class SqliteConnectionPool: """Thread-local connection manager for SQLite. @@ -149,28 +136,17 @@ def new_connection(self) -> SqliteConnection: connection = sqlite3.connect(**self._connection_parameters) try: - busy_timeout = self._connection_parameters.get("busy_timeout", SQLITE_BUSY_TIMEOUT) - connection.execute(f"PRAGMA busy_timeout = {busy_timeout}") - if self._enable_optimizations: if self._is_memory_db: connection.execute("PRAGMA journal_mode = MEMORY") connection.execute("PRAGMA synchronous = OFF") connection.execute("PRAGMA temp_store = MEMORY") - cache_size = self._connection_parameters.get("cache_size", SQLITE_MEMORY_CACHE_SIZE) - connection.execute(f"PRAGMA cache_size = {cache_size}") + connection.execute(f"PRAGMA cache_size = {SQLITE_MEMORY_CACHE_SIZE}") else: - current_mode = connection.execute("PRAGMA journal_mode").fetchone() - current_mode_str = str(current_mode[0]).lower() if current_mode else "" - if current_mode_str != "wal": - _enable_wal(connection) + _enable_wal(connection) connection.execute("PRAGMA synchronous = NORMAL") - cache_size = self._connection_parameters.get("cache_size", SQLITE_DISK_CACHE_SIZE) - connection.execute(f"PRAGMA cache_size = {cache_size}") - mmap_size = self._connection_parameters.get("mmap_size", SQLITE_MMAP_SIZE) - connection.execute(f"PRAGMA mmap_size = {mmap_size}") - connection.execute("PRAGMA temp_store = MEMORY") - connection.execute(f"PRAGMA journal_size_limit = {SQLITE_JOURNAL_SIZE_LIMIT}") + + connection.execute(f"PRAGMA busy_timeout = {SQLITE_BUSY_TIMEOUT}") if self._enable_foreign_keys: connection.execute("PRAGMA foreign_keys = ON") @@ -378,7 +354,7 @@ def _apply_runtime_setup(connection: SqliteConnection, runtime_setup: "dict[str, function_config["name"], function_config["narg"], function_config["func"], - deterministic=function_config.get("deterministic", True), + deterministic=function_config.get("deterministic", False), ) for aggregate_config in runtime_setup.get("custom_aggregates", ()): From 66e75032a4c56fcae770381a01c7bc365f273f60 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 02:18:12 +0000 Subject: [PATCH 03/12] fix(sqlite): guard autocommit attribute lookup and exclude driver from mypyc --- pyproject.toml | 1 - sqlspec/adapters/sqlite/driver.py | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index d9b4b0e80..6158692aa 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -292,7 +292,6 @@ include = [ "sqlspec/adapters/mysql_common.py", # Shared MySQL-family adapter helpers "sqlspec/adapters/**/core.py", # Adapter compiled helpers "sqlspec/adapters/**/type_converter.py", # All adapters type converters - "sqlspec/adapters/sqlite/driver.py", # SQLite synchronous driver "sqlspec/adapters/oracledb/_param_types.py", # Slot-based LOB/JSON parameter wrappers "sqlspec/adapters/oracledb/_json_handlers.py", # Native JSON inputtypehandler / outputtypehandler chain "sqlspec/adapters/oracledb/_uuid_handlers.py", # UUID ↔ RAW(16) inputtypehandler / outputtypehandler chain diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index 3141e41f6..f62334938 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -248,7 +248,7 @@ def _in_autocommit_mode(self) -> bool: no-ops, so a manually started transaction has to be ended with an explicit statement. """ - return SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT and self.connection.autocommit is True + return SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT and getattr(self.connection, "autocommit", None) is True def commit(self) -> None: """Commit the current transaction. From 3680bb21319f95cffa3c99517c4de1c35c527639 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 17:18:02 +0000 Subject: [PATCH 04/12] fix(sqlite): restore sqlite driver mypyc compilation and clean up helper imports --- pyproject.toml | 1 + sqlspec/adapters/aiosqlite/adk/store.py | 2 +- sqlspec/adapters/aiosqlite/config.py | 75 ++----------------- sqlspec/adapters/aiosqlite/core.py | 63 +++++++++++++++- sqlspec/adapters/aiosqlite/events/store.py | 3 +- sqlspec/adapters/aiosqlite/litestar/store.py | 2 +- sqlspec/adapters/sqlite/adk/store.py | 3 +- sqlspec/adapters/sqlite/config.py | 77 ++------------------ sqlspec/adapters/sqlite/core.py | 59 ++++++++++++++- sqlspec/adapters/sqlite/events/store.py | 3 +- sqlspec/adapters/sqlite/litestar/store.py | 3 +- sqlspec/adapters/sqlite/pool.py | 6 +- tests/unit/adapters/test_sqlite/test_pool.py | 9 ++- tools/scripts/mypyc_inventory.py | 4 +- 14 files changed, 154 insertions(+), 156 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 6158692aa..d9b4b0e80 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -292,6 +292,7 @@ include = [ "sqlspec/adapters/mysql_common.py", # Shared MySQL-family adapter helpers "sqlspec/adapters/**/core.py", # Adapter compiled helpers "sqlspec/adapters/**/type_converter.py", # All adapters type converters + "sqlspec/adapters/sqlite/driver.py", # SQLite synchronous driver "sqlspec/adapters/oracledb/_param_types.py", # Slot-based LOB/JSON parameter wrappers "sqlspec/adapters/oracledb/_json_handlers.py", # Native JSON inputtypehandler / outputtypehandler chain "sqlspec/adapters/oracledb/_uuid_handlers.py", # UUID ↔ RAW(16) inputtypehandler / outputtypehandler chain diff --git a/sqlspec/adapters/aiosqlite/adk/store.py b/sqlspec/adapters/aiosqlite/adk/store.py index ced1d772b..9d38ce2f7 100644 --- a/sqlspec/adapters/aiosqlite/adk/store.py +++ b/sqlspec/adapters/aiosqlite/adk/store.py @@ -8,7 +8,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 -from sqlspec.adapters.aiosqlite.config import render_pragmas +from sqlspec.adapters.aiosqlite.core import render_pragmas from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options diff --git a/sqlspec/adapters/aiosqlite/config.py b/sqlspec/adapters/aiosqlite/config.py index 2fbd1b382..5741c9e7b 100644 --- a/sqlspec/adapters/aiosqlite/config.py +++ b/sqlspec/adapters/aiosqlite/config.py @@ -1,7 +1,5 @@ """Aiosqlite database configuration.""" -import re -from collections.abc import Mapping from os import PathLike from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast @@ -14,7 +12,12 @@ AiosqliteCursor, AiosqliteSessionContext, ) -from sqlspec.adapters.aiosqlite.core import apply_driver_features, build_connection_config, default_statement_config +from sqlspec.adapters.aiosqlite.core import ( + apply_driver_features, + build_connection_config, + default_statement_config, + render_pragmas, +) from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver, AiosqliteExceptionHandler from sqlspec.adapters.aiosqlite.pool import ( AiosqliteConnectionPool, @@ -31,7 +34,7 @@ from sqlspec.utils.uuids import uuid4 if TYPE_CHECKING: - from collections.abc import Awaitable, Callable, Sequence + from collections.abc import Awaitable, Callable, Mapping, Sequence from types import TracebackType from sqlspec.core import StatementConfig @@ -45,9 +48,6 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", - "apply_extension_pragmas", - "extension_pragma_statements", - "render_pragmas", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -171,8 +171,6 @@ class AiosqliteDriverFeatures(TypedDict): extensions: "NotRequired[Sequence[str]]" -_PRAGMA_NAME_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") -_PRAGMA_VALUE_PATTERN = re.compile(r"^[A-Za-z0-9_.\-]+$") _ROW_FACTORY_LITERALS = frozenset({"dict", "row", "tuple"}) _RUNTIME_FEATURE_KEYS = ( "authorizer_callback", @@ -187,12 +185,6 @@ class AiosqliteDriverFeatures(TypedDict): "text_factory", "trace_callback", ) -_EXTENSION_PRAGMA_PROFILE = ( - "PRAGMA foreign_keys = ON", - "PRAGMA cache_size = -64000", - "PRAGMA mmap_size = 30000000", - "PRAGMA journal_size_limit = 67108864", -) class _AiosqliteSessionFactory(AsyncPoolSessionFactory): @@ -443,59 +435,6 @@ async def _close_pool(self) -> None: self.connection_instance = None -def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": - extension_config = cast("dict[str, Any]", config.extension_config) - settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) - profile = settings.get("pragma_profile", False) - if not isinstance(profile, bool): - msg = f"extension_config['{extension_name}']['pragma_profile'] must be a boolean" - raise ImproperConfigurationError(msg) - statements: list[str] = list(_EXTENSION_PRAGMA_PROFILE) if profile else [] - overrides = settings.get("pragma_overrides") - if overrides is None: - return tuple(statements) - if not isinstance(overrides, Mapping): - msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" - raise ImproperConfigurationError(msg) - try: - statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) - except ImproperConfigurationError as exc: - msg = str(exc).replace( - "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" - ) - raise ImproperConfigurationError(msg) from exc - return tuple(statements) - - -async def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: - for statement in statements: - await connection.execute(statement) - - -def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": - rendered: list[tuple[str, str]] = [] - for pragma_name, pragma_value in pragmas.items(): - if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: - msg = f"Invalid PRAGMA name in driver_features['pragmas']: {pragma_name!r}" - raise ImproperConfigurationError(msg) - if isinstance(pragma_value, bool): - rendered_value = "1" if pragma_value else "0" - elif isinstance(pragma_value, int): - rendered_value = str(pragma_value) - elif isinstance(pragma_value, str) and _PRAGMA_VALUE_PATTERN.match(pragma_value) is not None: - rendered_value = pragma_value - else: - msg = f"Invalid PRAGMA value for {pragma_name!r} in driver_features['pragmas']: {pragma_value!r}" - raise ImproperConfigurationError(msg) - rendered.append((pragma_name, rendered_value)) - return rendered - - -_extension_pragma_statements = extension_pragma_statements -_apply_extension_pragmas = apply_extension_pragmas -_render_pragmas = render_pragmas - - def _validate_entries(entries: Any, required_keys: "tuple[str, ...]", feature_name: str) -> None: for entry in entries: for required_key in required_keys: diff --git a/sqlspec/adapters/aiosqlite/core.py b/sqlspec/adapters/aiosqlite/core.py index 94465842f..7f449e7fa 100644 --- a/sqlspec/adapters/aiosqlite/core.py +++ b/sqlspec/adapters/aiosqlite/core.py @@ -1,7 +1,9 @@ """AIOSQLite adapter compiled helpers.""" import contextlib +import re import sys +from collections.abc import Mapping from datetime import date, datetime from decimal import Decimal from typing import TYPE_CHECKING, Any, TypeVar, cast @@ -33,7 +35,7 @@ from sqlspec.utils.type_guards import has_lastrowid, has_rowcount, has_sqlite_error if TYPE_CHECKING: - from collections.abc import Awaitable, Callable, Mapping, Sequence + from collections.abc import Awaitable, Callable, Sequence from sqlspec.adapters.aiosqlite._typing import AiosqliteConnection from sqlspec.core.compiler import OperationType @@ -43,6 +45,7 @@ __all__ = ( "AiosqliteStreamSource", "apply_driver_features", + "apply_extension_pragmas", "build_connection_config", "build_insert_statement", "build_profile", @@ -57,10 +60,12 @@ "execute_fetchall_with_metadata", "execute_many_on_worker_thread", "execute_script_on_worker_thread", + "extension_pragma_statements", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", "normalize_lastrowid", + "render_pragmas", "require_python_version", "resolve_lastrowid", "resolve_rowcount", @@ -70,6 +75,14 @@ _TIME_TO_ISO = time_iso_convert _DECIMAL_TO_STRING = build_decimal_converter(mode="string") +_PRAGMA_NAME_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +_PRAGMA_VALUE_PATTERN = re.compile(r"^[A-Za-z0-9_.\-]+$") +_EXTENSION_PRAGMA_PROFILE = ( + "PRAGMA foreign_keys = ON", + "PRAGMA cache_size = -64000", + "PRAGMA mmap_size = 30000000", + "PRAGMA journal_size_limit = 67108864", +) SQLITE_CONSTRAINT_UNIQUE_CODE = 2067 SQLITE_CONSTRAINT_PRIMARYKEY_CODE = 1555 @@ -484,6 +497,54 @@ def apply_driver_features( return statement_config, features +def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": + extension_config = cast("dict[str, Any]", config.extension_config) + settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) + profile = settings.get("pragma_profile", False) + if not isinstance(profile, bool): + msg = f"extension_config['{extension_name}']['pragma_profile'] must be a boolean" + raise ImproperConfigurationError(msg) + statements: list[str] = list(_EXTENSION_PRAGMA_PROFILE) if profile else [] + overrides = settings.get("pragma_overrides") + if overrides is None: + return tuple(statements) + if not isinstance(overrides, Mapping): + msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" + raise ImproperConfigurationError(msg) + try: + statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) + except ImproperConfigurationError as exc: + msg = str(exc).replace( + "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" + ) + raise ImproperConfigurationError(msg) from exc + return tuple(statements) + + +async def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: + for statement in statements: + await connection.execute(statement) + + +def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": + rendered: list[tuple[str, str]] = [] + for pragma_name, pragma_value in pragmas.items(): + if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: + msg = f"Invalid PRAGMA name in driver_features['pragmas']: {pragma_name!r}" + raise ImproperConfigurationError(msg) + if isinstance(pragma_value, bool): + rendered_value = "1" if pragma_value else "0" + elif isinstance(pragma_value, int): + rendered_value = str(pragma_value) + elif isinstance(pragma_value, str) and _PRAGMA_VALUE_PATTERN.match(pragma_value) is not None: + rendered_value = pragma_value + else: + msg = f"Invalid PRAGMA value for {pragma_name!r} in driver_features['pragmas']: {pragma_value!r}" + raise ImproperConfigurationError(msg) + rendered.append((pragma_name, rendered_value)) + return rendered + + def _execute_fetchall_with_metadata( connection: "AiosqliteConnection", sql: str, diff --git a/sqlspec/adapters/aiosqlite/events/store.py b/sqlspec/adapters/aiosqlite/events/store.py index dedff444a..90a842092 100644 --- a/sqlspec/adapters/aiosqlite/events/store.py +++ b/sqlspec/adapters/aiosqlite/events/store.py @@ -4,7 +4,8 @@ from typing_extensions import NotRequired -from sqlspec.adapters.aiosqlite.config import AiosqliteConfig, apply_extension_pragmas, extension_pragma_statements +from sqlspec.adapters.aiosqlite.config import AiosqliteConfig +from sqlspec.adapters.aiosqlite.core import apply_extension_pragmas, extension_pragma_statements from sqlspec.config import EventsConfig from sqlspec.extensions.events import BaseEventQueueStore diff --git a/sqlspec/adapters/aiosqlite/litestar/store.py b/sqlspec/adapters/aiosqlite/litestar/store.py index aece8c9e6..ddffcafba 100644 --- a/sqlspec/adapters/aiosqlite/litestar/store.py +++ b/sqlspec/adapters/aiosqlite/litestar/store.py @@ -5,7 +5,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.aiosqlite.config import apply_extension_pragmas, extension_pragma_statements +from sqlspec.adapters.aiosqlite.core import apply_extension_pragmas, extension_pragma_statements from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore diff --git a/sqlspec/adapters/sqlite/adk/store.py b/sqlspec/adapters/sqlite/adk/store.py index 1f39a1e3a..4c22f5ffd 100644 --- a/sqlspec/adapters/sqlite/adk/store.py +++ b/sqlspec/adapters/sqlite/adk/store.py @@ -8,8 +8,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 -from sqlspec.adapters.sqlite.config import render_pragmas -from sqlspec.adapters.sqlite.core import end_transaction +from sqlspec.adapters.sqlite.core import end_transaction, render_pragmas from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options diff --git a/sqlspec/adapters/sqlite/config.py b/sqlspec/adapters/sqlite/config.py index b31efff57..2916ab93b 100644 --- a/sqlspec/adapters/sqlite/config.py +++ b/sqlspec/adapters/sqlite/config.py @@ -1,9 +1,7 @@ """SQLite database configuration with thread-local connections.""" -import re -from collections.abc import Mapping from os import PathLike -from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, TypedDict from typing_extensions import NotRequired @@ -13,7 +11,12 @@ SqliteCursor, SqliteSessionContext, ) -from sqlspec.adapters.sqlite.core import apply_driver_features, build_connection_config, default_statement_config +from sqlspec.adapters.sqlite.core import ( + apply_driver_features, + build_connection_config, + default_statement_config, + render_pragmas, +) from sqlspec.adapters.sqlite.driver import SqliteDriver, SqliteExceptionHandler from sqlspec.adapters.sqlite.pool import SqliteConnectionPool from sqlspec.adapters.sqlite.type_converter import register_type_handlers @@ -25,7 +28,7 @@ from sqlspec.utils.uuids import uuid4 if TYPE_CHECKING: - from collections.abc import Callable, Sequence + from collections.abc import Callable, Mapping, Sequence from sqlspec.core import StatementConfig from sqlspec.observability import ObservabilityConfig @@ -38,9 +41,6 @@ "SqliteDriverFeatures", "SqliteFunctionConfig", "SqliteWindowFunctionConfig", - "apply_extension_pragmas", - "extension_pragma_statements", - "render_pragmas", ) logger = get_logger("sqlspec.adapters.sqlite") @@ -164,8 +164,6 @@ class SqliteDriverFeatures(TypedDict): extensions: "NotRequired[Sequence[str]]" -_PRAGMA_NAME_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") -_PRAGMA_VALUE_PATTERN = re.compile(r"^[A-Za-z0-9_.\-]+$") _ROW_FACTORY_LITERALS = frozenset({"dict", "row", "tuple"}) _RUNTIME_FEATURE_KEYS = ( "authorizer_callback", @@ -181,12 +179,6 @@ class SqliteDriverFeatures(TypedDict): "text_factory", "trace_callback", ) -_EXTENSION_PRAGMA_PROFILE = ( - "PRAGMA foreign_keys = ON", - "PRAGMA cache_size = -64000", - "PRAGMA mmap_size = 30000000", - "PRAGMA journal_size_limit = 67108864", -) class SqliteConnectionContext(SyncPoolConnectionContext): @@ -376,59 +368,6 @@ def _close_pool(self) -> None: self.connection_instance.close() -def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": - extension_config = cast("dict[str, Any]", config.extension_config) - settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) - profile = settings.get("pragma_profile", False) - if not isinstance(profile, bool): - msg = f"extension_config['{extension_name}']['pragma_profile'] must be a boolean" - raise ImproperConfigurationError(msg) - statements: list[str] = list(_EXTENSION_PRAGMA_PROFILE) if profile else [] - overrides = settings.get("pragma_overrides") - if overrides is None: - return tuple(statements) - if not isinstance(overrides, Mapping): - msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" - raise ImproperConfigurationError(msg) - try: - statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) - except ImproperConfigurationError as exc: - msg = str(exc).replace( - "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" - ) - raise ImproperConfigurationError(msg) from exc - return tuple(statements) - - -def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: - for statement in statements: - connection.execute(statement) - - -def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": - rendered: list[tuple[str, str]] = [] - for pragma_name, pragma_value in pragmas.items(): - if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: - msg = f"Invalid PRAGMA name in driver_features['pragmas']: {pragma_name!r}" - raise ImproperConfigurationError(msg) - if isinstance(pragma_value, bool): - rendered_value = "1" if pragma_value else "0" - elif isinstance(pragma_value, int): - rendered_value = str(pragma_value) - elif isinstance(pragma_value, str) and _PRAGMA_VALUE_PATTERN.match(pragma_value) is not None: - rendered_value = pragma_value - else: - msg = f"Invalid PRAGMA value for {pragma_name!r} in driver_features['pragmas']: {pragma_value!r}" - raise ImproperConfigurationError(msg) - rendered.append((pragma_name, rendered_value)) - return rendered - - -_extension_pragma_statements = extension_pragma_statements -_apply_extension_pragmas = apply_extension_pragmas -_render_pragmas = render_pragmas - - def _validate_entries(entries: Any, required_keys: "tuple[str, ...]", feature_name: str) -> None: for entry in entries: for required_key in required_keys: diff --git a/sqlspec/adapters/sqlite/core.py b/sqlspec/adapters/sqlite/core.py index dacbc498a..883170a50 100644 --- a/sqlspec/adapters/sqlite/core.py +++ b/sqlspec/adapters/sqlite/core.py @@ -1,6 +1,7 @@ """SQLite adapter compiled helpers.""" import contextlib +import re import sys from collections.abc import Mapping from datetime import date, datetime @@ -43,6 +44,7 @@ "SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", "SqliteStreamSource", "apply_driver_features", + "apply_extension_pragmas", "build_connection_config", "build_insert_statement", "build_profile", @@ -52,10 +54,12 @@ "default_statement_config", "driver_profile", "end_transaction", + "extension_pragma_statements", "format_identifier", "normalize_execute_many_parameters", "normalize_execute_parameters", "normalize_lastrowid", + "render_pragmas", "require_python_version", "resolve_lastrowid", "resolve_rowcount", @@ -80,6 +84,14 @@ SQLITE_TABLE_LIST_MIN_COLUMNS = 5 SQLITE_TABLE_INFO_MIN_COLUMNS = 2 SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") +_PRAGMA_NAME_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +_PRAGMA_VALUE_PATTERN = re.compile(r"^[A-Za-z0-9_.\-]+$") +_EXTENSION_PRAGMA_PROFILE = ( + "PRAGMA foreign_keys = ON", + "PRAGMA cache_size = -64000", + "PRAGMA mmap_size = 30000000", + "PRAGMA journal_size_limit = 67108864", +) def end_transaction( @@ -109,7 +121,52 @@ def end_transaction( connection.rollback() -_end_transaction = end_transaction +def extension_pragma_statements(config: Any, extension_name: str) -> "tuple[str, ...]": + extension_config = cast("dict[str, Any]", config.extension_config) + settings = cast("dict[str, Any]", extension_config.get(extension_name, {})) + profile = settings.get("pragma_profile", False) + if not isinstance(profile, bool): + msg = f"extension_config['{extension_name}']['pragma_profile'] must be a boolean" + raise ImproperConfigurationError(msg) + statements: list[str] = list(_EXTENSION_PRAGMA_PROFILE) if profile else [] + overrides = settings.get("pragma_overrides") + if overrides is None: + return tuple(statements) + if not isinstance(overrides, Mapping): + msg = f"extension_config['{extension_name}']['pragma_overrides'] must be a mapping of PRAGMA names to values" + raise ImproperConfigurationError(msg) + try: + statements.extend(f"PRAGMA {name} = {value}" for name, value in render_pragmas(overrides)) + except ImproperConfigurationError as exc: + msg = str(exc).replace( + "driver_features['pragmas']", f"extension_config['{extension_name}']['pragma_overrides']" + ) + raise ImproperConfigurationError(msg) from exc + return tuple(statements) + + +def apply_extension_pragmas(connection: Any, statements: "tuple[str, ...]") -> None: + for statement in statements: + connection.execute(statement) + + +def render_pragmas(pragmas: "Mapping[str, Any]") -> "list[tuple[str, str]]": + rendered: list[tuple[str, str]] = [] + for pragma_name, pragma_value in pragmas.items(): + if not isinstance(pragma_name, str) or _PRAGMA_NAME_PATTERN.match(pragma_name) is None: + msg = f"Invalid PRAGMA name in driver_features['pragmas']: {pragma_name!r}" + raise ImproperConfigurationError(msg) + if isinstance(pragma_value, bool): + rendered_value = "1" if pragma_value else "0" + elif isinstance(pragma_value, int): + rendered_value = str(pragma_value) + elif isinstance(pragma_value, str) and _PRAGMA_VALUE_PATTERN.match(pragma_value) is not None: + rendered_value = pragma_value + else: + msg = f"Invalid PRAGMA value for {pragma_name!r} in driver_features['pragmas']: {pragma_value!r}" + raise ImproperConfigurationError(msg) + rendered.append((pragma_name, rendered_value)) + return rendered _TIME_TO_ISO = time_iso_convert diff --git a/sqlspec/adapters/sqlite/events/store.py b/sqlspec/adapters/sqlite/events/store.py index ef2da88e3..0bba39693 100644 --- a/sqlspec/adapters/sqlite/events/store.py +++ b/sqlspec/adapters/sqlite/events/store.py @@ -4,7 +4,8 @@ from typing_extensions import NotRequired -from sqlspec.adapters.sqlite.config import SqliteConfig, apply_extension_pragmas, extension_pragma_statements +from sqlspec.adapters.sqlite.config import SqliteConfig +from sqlspec.adapters.sqlite.core import apply_extension_pragmas, extension_pragma_statements from sqlspec.config import EventsConfig from sqlspec.extensions.events import BaseEventQueueStore diff --git a/sqlspec/adapters/sqlite/litestar/store.py b/sqlspec/adapters/sqlite/litestar/store.py index b8875536d..b63892b0b 100644 --- a/sqlspec/adapters/sqlite/litestar/store.py +++ b/sqlspec/adapters/sqlite/litestar/store.py @@ -5,8 +5,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.sqlite.config import apply_extension_pragmas, extension_pragma_statements -from sqlspec.adapters.sqlite.core import end_transaction +from sqlspec.adapters.sqlite.core import apply_extension_pragmas, end_transaction, extension_pragma_statements from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore from sqlspec.utils.sync_tools import async_ diff --git a/sqlspec/adapters/sqlite/pool.py b/sqlspec/adapters/sqlite/pool.py index aae34d70d..7a32fd084 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -9,7 +9,7 @@ from sqlspec.adapters.sqlite._typing import SqliteConnection from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 -from sqlspec.adapters.sqlite.core import end_transaction as _end_transaction +from sqlspec.adapters.sqlite.core import end_transaction from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -273,11 +273,11 @@ def get_connection(self) -> "Generator[SqliteConnection, None, None]": yield connection except Exception: with contextlib.suppress(Exception): - _end_transaction(connection, commit=False) + end_transaction(connection, commit=False) raise else: with contextlib.suppress(Exception): - _end_transaction(connection, commit=True) + end_transaction(connection, commit=True) def close(self) -> None: """Close every connection this pool opened, on any thread.""" diff --git a/tests/unit/adapters/test_sqlite/test_pool.py b/tests/unit/adapters/test_sqlite/test_pool.py index 289c34f61..19729d5c7 100644 --- a/tests/unit/adapters/test_sqlite/test_pool.py +++ b/tests/unit/adapters/test_sqlite/test_pool.py @@ -7,7 +7,8 @@ import pytest -from sqlspec.adapters.sqlite.pool import SqliteConnectionPool, _end_transaction +from sqlspec.adapters.sqlite.core import end_transaction +from sqlspec.adapters.sqlite.pool import SqliteConnectionPool if TYPE_CHECKING: from sqlspec.adapters.sqlite._typing import SqliteConnection @@ -247,7 +248,7 @@ def test_pool_ends_autocommit_transactions_with_an_explicit_statement(commit: bo """The pool's own commit is a no-op in autocommit mode, exactly like the driver's.""" connection = _AutocommitConnection() - _end_transaction(cast("Any", connection), commit=commit, supports_autocommit=True) + end_transaction(cast("Any", connection), commit=commit, supports_autocommit=True) assert connection.statements == [statement] assert connection.commit_calls == 0 @@ -259,7 +260,7 @@ def test_pool_uses_the_dbapi_methods_outside_autocommit_mode() -> None: connection = _AutocommitConnection() connection.autocommit = False - _end_transaction(cast("Any", connection), commit=True, supports_autocommit=True) + end_transaction(cast("Any", connection), commit=True, supports_autocommit=True) assert connection.statements == [] assert connection.commit_calls == 1 @@ -269,7 +270,7 @@ def test_pool_skips_ending_a_transaction_that_is_not_open() -> None: connection = _AutocommitConnection() connection.in_transaction = False - _end_transaction(cast("Any", connection), commit=True, supports_autocommit=True) + end_transaction(cast("Any", connection), commit=True, supports_autocommit=True) assert connection.statements == [] assert connection.commit_calls == 0 diff --git a/tools/scripts/mypyc_inventory.py b/tools/scripts/mypyc_inventory.py index e618f1cdd..d1fa03953 100644 --- a/tools/scripts/mypyc_inventory.py +++ b/tools/scripts/mypyc_inventory.py @@ -234,8 +234,8 @@ "reason": "Async pool runtime compiles without adapter driver exception-handler subclasses.", }, "sqlspec/adapters/sqlite/driver.py": { - "classification": "prove_separately", - "reason": "Native compiled driver construction segfaulted in installed-wheel SqliteConfig session smoke; keep interpreted pending driver layout work.", + "classification": "compile_now", + "reason": "Synchronous SQLite driver compiles cleanly and passes installed-wheel session smoke.", }, "sqlspec/adapters/aiosqlite/driver.py": { "classification": "prove_separately", From 0540862500307aca7910a9e1ea5df142ccbf9425 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 18:13:45 +0000 Subject: [PATCH 05/12] fix(sqlite): import compiled driver submodules directly for mypyc wheel build --- sqlspec/adapters/sqlite/driver.py | 35 ++++++++++++++++++++----------- 1 file changed, 23 insertions(+), 12 deletions(-) diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index f62334938..75b26ceaf 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -2,6 +2,8 @@ from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast +from mypy_extensions import mypyc_attr + from sqlspec.adapters.sqlite._typing import SqliteCursor, SqliteSessionContext from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 from sqlspec.adapters.sqlite.core import ( @@ -19,27 +21,34 @@ resolve_rowcount, ) from sqlspec.adapters.sqlite.data_dictionary import SqliteDataDictionary -from sqlspec.core import ArrowResult, ParameterStyle, TypedParameter, get_cache_config, register_driver_profile -from sqlspec.core.result import DMLResult -from sqlspec.driver import ( - BaseSyncExceptionHandler, - SyncDriverAdapterBase, - SyncRowStream, +from sqlspec.core.cache import get_cache_config +from sqlspec.core.parameters._registry import register_driver_profile +from sqlspec.core.parameters._types import ParameterStyle, TypedParameter +from sqlspec.core.result._base import ArrowResult, DMLResult +from sqlspec.driver._common import ( + CachedQuery, + ExecutionResult, parameter_value_needs_processing, type_coercion_fallbacks, ) +from sqlspec.driver._exception_handler import BaseSyncExceptionHandler +from sqlspec.driver._stream import SyncRowStream +from sqlspec.driver._sync import SyncDriverAdapterBase from sqlspec.exceptions import SQLSpecError from sqlspec.utils.type_guards import resolve_row_format if TYPE_CHECKING: from collections.abc import Sequence + from sqlglot.dialects.dialect import DialectType + from sqlspec.adapters.sqlite._typing import SqliteConnection - from sqlspec.builder import QueryBuilder - from sqlspec.core import SQL, SQLResult, Statement, StatementConfig, StatementFilter + from sqlspec.builder._base import QueryBuilder from sqlspec.core.compiler import OperationType - from sqlspec.driver import CachedQuery, ExecutionResult - from sqlspec.storage import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry + from sqlspec.core.filters import StatementFilter + from sqlspec.core.result._base import SQLResult + from sqlspec.core.statement import SQL, Statement, StatementConfig + from sqlspec.storage.pipeline import StorageBridgeJob, StorageDestination, StorageFormat, StorageTelemetry from sqlspec.typing import StatementParameters __all__ = ("SqliteCursor", "SqliteDriver", "SqliteExceptionHandler", "SqliteSessionContext") @@ -48,6 +57,7 @@ _BATCH_SAMPLE_THRESHOLD: Final = 100 +@mypyc_attr(allow_interpreted_subclasses=True) class SqliteExceptionHandler(BaseSyncExceptionHandler): """Context manager for handling SQLite database exceptions. @@ -70,6 +80,7 @@ def _handle_exception(self, exc_type: "type[BaseException] | None", exc_val: "Ba return False +@mypyc_attr(allow_interpreted_subclasses=True) class SqliteDriver(SyncDriverAdapterBase): """SQLite driver implementation. @@ -78,7 +89,7 @@ class SqliteDriver(SyncDriverAdapterBase): """ __slots__ = ("_data_dictionary", "_rowid_target_cache") - dialect = "sqlite" + dialect: "DialectType | None" = "sqlite" def __init__( self, @@ -359,7 +370,7 @@ def load_from_arrow( cursor.execute(statement) for batch in arrow_table.to_batches(max_chunksize=batch_size): pydict = batch.to_pydict() - records = list(zip(*(pydict[col] for col in columns), strict=False)) + records = list(zip(*[pydict[col] for col in columns], strict=False)) if records: prepared_records = ( self.prepare_driver_parameters(records, self.statement_config, is_many=True) From 40c7e37bb3fc43e027f1e0d41635a6a97e869ca5 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sat, 26 Sep 2026 20:30:24 +0000 Subject: [PATCH 06/12] refactor(sqlite): simplify _typing.py by removing redundant TYPE_CHECKING splits --- sqlspec/adapters/aiosqlite/_typing.py | 21 +++++++------------ sqlspec/adapters/sqlite/_typing.py | 21 +++++++------------ .../adapters/test_aiosqlite/test_config.py | 2 +- 3 files changed, 16 insertions(+), 28 deletions(-) diff --git a/sqlspec/adapters/aiosqlite/_typing.py b/sqlspec/adapters/aiosqlite/_typing.py index 711cd0a26..e908351e6 100644 --- a/sqlspec/adapters/aiosqlite/_typing.py +++ b/sqlspec/adapters/aiosqlite/_typing.py @@ -5,35 +5,28 @@ """ import contextlib -import sqlite3 import sqlite3 as aiosqlite_sqlite_module -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, TypeAlias -import aiosqlite import aiosqlite as aiosqlite_module -from typing_extensions import TypeAliasType +from aiosqlite import Connection as AiosqliteConnection +from aiosqlite import Cursor as AiosqliteRawCursor +from aiosqlite import Error as AiosqliteError + +AiosqliteConnectionFactory: TypeAlias = type[aiosqlite_sqlite_module.Connection] if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver from sqlspec.core import StatementConfig - AiosqliteConnection: TypeAlias = aiosqlite.Connection - AiosqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] - AiosqliteRawCursor: TypeAlias = aiosqlite.Cursor - -if not TYPE_CHECKING: - AiosqliteConnection = aiosqlite.Connection - AiosqliteConnectionFactory = TypeAliasType("AiosqliteConnectionFactory", type[sqlite3.Connection]) - AiosqliteRawCursor = aiosqlite.Cursor - __all__ = ( "AiosqliteConnection", "AiosqliteConnectionFactory", "AiosqliteCursor", + "AiosqliteError", "AiosqliteRawCursor", "AiosqliteSessionContext", "aiosqlite_module", diff --git a/sqlspec/adapters/sqlite/_typing.py b/sqlspec/adapters/sqlite/_typing.py index 25dfe5676..7a974adad 100644 --- a/sqlspec/adapters/sqlite/_typing.py +++ b/sqlspec/adapters/sqlite/_typing.py @@ -5,33 +5,28 @@ """ import contextlib -import sqlite3 import sqlite3 as sqlite_module -from typing import TYPE_CHECKING, Any +from sqlite3 import Connection as SqliteConnection +from sqlite3 import Cursor as SqliteRawCursor +from sqlite3 import Error as SqliteError +from sqlite3 import OperationalError as SqliteOperationalError +from typing import TYPE_CHECKING, Any, TypeAlias -_SqliteConnection = sqlite3.Connection +SqliteConnectionFactory: TypeAlias = type[SqliteConnection] if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType - from typing import TypeAlias from sqlspec.adapters.sqlite.driver import SqliteDriver from sqlspec.core import StatementConfig - SqliteConnection: TypeAlias = _SqliteConnection - SqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] - SqliteRawCursor: TypeAlias = sqlite3.Cursor - -if not TYPE_CHECKING: - SqliteConnection = _SqliteConnection - SqliteConnectionFactory = type[sqlite3.Connection] - SqliteRawCursor = sqlite3.Cursor - __all__ = ( "SqliteConnection", "SqliteConnectionFactory", "SqliteCursor", + "SqliteError", + "SqliteOperationalError", "SqliteRawCursor", "SqliteSessionContext", "sqlite_module", diff --git a/tests/unit/adapters/test_aiosqlite/test_config.py b/tests/unit/adapters/test_aiosqlite/test_config.py index 7cd6c43d6..4b13f2eb7 100644 --- a/tests/unit/adapters/test_aiosqlite/test_config.py +++ b/tests/unit/adapters/test_aiosqlite/test_config.py @@ -21,7 +21,7 @@ class CustomConnection(sqlite3.Connection): def _annotation_contains(annotation: object, expected: object) -> bool: """Return whether an annotation tree contains the expected object.""" - if annotation is expected: + if annotation is expected or annotation == expected: return True return any(_annotation_contains(arg, expected) for arg in get_args(annotation)) From 77cebd76514d418cef8f8f649e0a63f7cb27d22e Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 15:19:35 +0000 Subject: [PATCH 07/12] fix(sqlite): close cursors on fast paths, validate all batch rows, and add aiosqlite end_transaction --- sqlspec/adapters/aiosqlite/__init__.py | 2 + sqlspec/adapters/aiosqlite/adk/store.py | 89 ++++++++++--------- sqlspec/adapters/aiosqlite/config.py | 19 ++++ sqlspec/adapters/aiosqlite/core.py | 67 ++++++++++++++ sqlspec/adapters/aiosqlite/driver.py | 34 ++++--- sqlspec/adapters/aiosqlite/litestar/store.py | 19 ++-- sqlspec/adapters/aiosqlite/pool.py | 26 ++++-- sqlspec/adapters/sqlite/__init__.py | 2 + sqlspec/adapters/sqlite/driver.py | 33 +++---- .../adapters/_shared/_driver_type_system.py | 2 + .../adapters/test_aiosqlite/test_config.py | 5 ++ .../adapters/test_aiosqlite/test_driver.py | 71 +++++++++++++++ .../unit/adapters/test_sqlite/test_driver.py | 76 +++++++++++++++- 13 files changed, 359 insertions(+), 86 deletions(-) diff --git a/sqlspec/adapters/aiosqlite/__init__.py b/sqlspec/adapters/aiosqlite/__init__.py index 1f48983b9..daa39414d 100644 --- a/sqlspec/adapters/aiosqlite/__init__.py +++ b/sqlspec/adapters/aiosqlite/__init__.py @@ -7,6 +7,7 @@ AiosqliteDriverFeatures, AiosqliteFunctionConfig, AiosqlitePoolParams, + AiosqliteWindowFunctionConfig, ) from sqlspec.adapters.aiosqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver, AiosqliteExceptionHandler @@ -34,6 +35,7 @@ "AiosqlitePoolConnection", "AiosqlitePoolParams", "AiosqliteRawCursor", + "AiosqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/aiosqlite/adk/store.py b/sqlspec/adapters/aiosqlite/adk/store.py index 9d38ce2f7..f863a2c53 100644 --- a/sqlspec/adapters/aiosqlite/adk/store.py +++ b/sqlspec/adapters/aiosqlite/adk/store.py @@ -8,7 +8,7 @@ from typing_extensions import NotRequired from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 -from sqlspec.adapters.aiosqlite.core import render_pragmas +from sqlspec.adapters.aiosqlite.core import end_transaction, render_pragmas from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseAsyncADKStore, StoredEvent, StoredSession, normalize_session_list_options @@ -128,7 +128,7 @@ async def create_session( async with self._config.provide_connection() as conn: await self._apply_pragmas(conn) await conn.execute(sql, params) - await conn.commit() + await end_transaction(conn, commit=True) return StoredSession( id=session_id, app_name=app_name, user_id=user_id, state=state, create_time=now, update_time=now @@ -165,7 +165,7 @@ async def get_session( WHERE app_name = ? AND user_id = ? AND id = ? """ await conn.execute(update_sql, (_datetime_to_julian(datetime.now(timezone.utc)), *params)) - await conn.commit() + await end_transaction(conn, commit=True) cursor = await conn.execute(sql, params) row = await cursor.fetchone() @@ -206,7 +206,7 @@ async def update_session_state(self, app_name: str, user_id: str, session_id: st async with self._config.provide_connection() as conn: await self._apply_pragmas(conn) await conn.execute(sql, (state_json, now_julian, app_name, user_id, session_id)) - await conn.commit() + await end_transaction(conn, commit=True) async def list_sessions( self, @@ -274,7 +274,7 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> async with self._config.provide_connection() as conn: await self._apply_pragmas(conn) await conn.execute(sql, (app_name, user_id, session_id)) - await conn.commit() + await end_transaction(conn, commit=True) async def append_event(self, event_record: StoredEvent) -> None: """Append an event to a session. @@ -305,7 +305,7 @@ async def append_event(self, event_record: StoredEvent) -> None: event_data_json, ), ) - await conn.commit() + await end_transaction(conn, commit=True) async def append_event_and_update_state( self, @@ -391,13 +391,13 @@ async def append_event_and_update_state( if user_state is not None: await conn.execute(user_upsert_sql, (app_name, user_id, to_json(user_state), now_julian)) except Exception: - await conn.rollback() + await end_transaction(conn, commit=False) raise else: if row is None: - await conn.rollback() + await end_transaction(conn, commit=False) else: - await conn.commit() + await end_transaction(conn, commit=True) if row is None: msg = f"Session {session_id} not found during append_event_and_update_state." @@ -488,7 +488,7 @@ async def delete_expired_events(self, before: datetime, app_name: "str | None" = await self._apply_pragmas(conn) cursor = await conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - await conn.commit() + await end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -508,7 +508,7 @@ async def delete_idle_sessions(self, updated_before: datetime, app_name: "str | await self._apply_pragmas(conn) cursor = await conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - await conn.commit() + await end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -528,7 +528,7 @@ async def delete_idle_user_states(self, updated_before: datetime, app_name: "str await self._apply_pragmas(conn) cursor = await conn.execute(sql, tuple(params)) deleted_count = cursor.rowcount if cursor.rowcount and cursor.rowcount > 0 else 0 - await conn.commit() + await end_transaction(conn, commit=True) return deleted_count except sqlite3.OperationalError as exc: if SQLITE_TABLE_NOT_FOUND_ERROR in str(exc): @@ -582,7 +582,7 @@ async def upsert_app_state(self, app_name: str, state: "dict[str, Any]") -> None async with self._config.provide_connection() as conn: await self._apply_pragmas(conn) await conn.execute(sql, (app_name, to_json(state), _datetime_to_julian(datetime.now(timezone.utc)))) - await conn.commit() + await end_transaction(conn, commit=True) async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, Any]") -> None: """Insert or replace user-scoped state for an application user.""" @@ -599,7 +599,7 @@ async def upsert_user_state(self, app_name: str, user_id: str, state: "dict[str, await conn.execute( sql, (app_name, user_id, to_json(state), _datetime_to_julian(datetime.now(timezone.utc))) ) - await conn.commit() + await end_transaction(conn, commit=True) async def get_metadata(self, key: str) -> "str | None": """Return a value from the ADK internal metadata table.""" @@ -627,7 +627,7 @@ async def set_metadata(self, key: str, value: str) -> None: async with self._config.provide_connection() as conn: await self._apply_pragmas(conn) await conn.execute(sql, (key, value)) - await conn.commit() + await end_transaction(conn, commit=True) async def _apply_pragmas(self, connection: Any) -> None: """Apply PRAGMA optimization profile for this connection. @@ -799,20 +799,19 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " if not entries: return 0 - inserted_count = 0 async with self._config.provide_connection() as conn: - for entry in entries: - params: tuple[Any, ...] - scope = entry.get("scope", "user") - if self._owner_id_column_name: - sql = f""" - INSERT OR IGNORE INTO {self._memory_table} - (id, session_id, app_name, user_id, scope, event_id, author, - {self._owner_id_column_name}, timestamp, content_json, - content_text, metadata_json, inserted_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - """ - params = ( + params_list: list[tuple[Any, ...]] = [] + if self._owner_id_column_name: + sql = f""" + INSERT OR IGNORE INTO {self._memory_table} + (id, session_id, app_name, user_id, scope, event_id, author, + {self._owner_id_column_name}, timestamp, content_json, + content_text, metadata_json, inserted_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """ + for entry in entries: + scope = entry.get("scope", "user") + params_list.append(( entry["id"], entry["session_id"], entry["app_name"], @@ -826,15 +825,17 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["content_text"], to_json(entry["metadata_json"]), _datetime_to_julian(entry["inserted_at"]), - ) - else: - sql = f""" - INSERT OR IGNORE INTO {self._memory_table} - (id, session_id, app_name, user_id, scope, event_id, author, - timestamp, content_json, content_text, metadata_json, inserted_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - """ - params = ( + )) + else: + sql = f""" + INSERT OR IGNORE INTO {self._memory_table} + (id, session_id, app_name, user_id, scope, event_id, author, + timestamp, content_json, content_text, metadata_json, inserted_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """ + for entry in entries: + scope = entry.get("scope", "user") + params_list.append(( entry["id"], entry["session_id"], entry["app_name"], @@ -847,11 +848,13 @@ async def insert_memory_entries(self, entries: "list[StoredMemory]", owner_id: " entry["content_text"], to_json(entry["metadata_json"]), _datetime_to_julian(entry["inserted_at"]), - ) - cursor = await conn.execute(sql, params) - inserted_count += cursor.rowcount + )) + cursor = await conn.executemany(sql, params_list) + try: + inserted_count = cursor.rowcount if cursor.rowcount >= 0 else len(params_list) + finally: await cursor.close() - await conn.commit() + await end_transaction(conn, commit=True) return inserted_count async def search_entries( @@ -917,7 +920,7 @@ async def delete_entries_by_session(self, session_id: str) -> int: sql = f"DELETE FROM {self._memory_table} WHERE session_id = ?" async with self._config.provide_connection() as conn: cursor = await conn.execute(sql, (session_id,)) - await conn.commit() + await end_transaction(conn, commit=True) return cursor.rowcount async def delete_entries_older_than( @@ -940,7 +943,7 @@ async def delete_entries_older_than( async with self._config.provide_connection() as conn: cursor = await conn.execute(sql, tuple(params)) - await conn.commit() + await end_transaction(conn, commit=True) return cursor.rowcount async def _memory_table_ddl(self) -> str: diff --git a/sqlspec/adapters/aiosqlite/config.py b/sqlspec/adapters/aiosqlite/config.py index 5741c9e7b..d5c80338d 100644 --- a/sqlspec/adapters/aiosqlite/config.py +++ b/sqlspec/adapters/aiosqlite/config.py @@ -48,6 +48,7 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", + "AiosqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -109,6 +110,14 @@ class AiosqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" +class AiosqliteWindowFunctionConfig(TypedDict): + """User-defined aiosqlite window function registration.""" + + name: str + narg: int + window_class: "type[Any]" + + class AiosqliteDriverFeatures(TypedDict): """Aiosqlite driver feature configuration. @@ -138,6 +147,9 @@ class AiosqliteDriverFeatures(TypedDict): Each entry must include name and func. Callable values are plain sync callables. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. + custom_window_functions: Register user-defined aggregate window functions. + Each entry must include name, narg, and window_class. + default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation on the worker thread. trace_callback: sqlite3 trace hook run for executed statements on the worker thread. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -161,6 +173,8 @@ class AiosqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[AiosqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[AiosqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[AiosqliteAggregateConfig]]" + custom_window_functions: "NotRequired[Sequence[AiosqliteWindowFunctionConfig]]" + default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -177,6 +191,7 @@ class AiosqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -348,6 +363,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "AiosqliteFunctionConfig": AiosqliteFunctionConfig, "AiosqlitePoolParams": AiosqlitePoolParams, "AiosqliteSessionContext": AiosqliteSessionContext, + "AiosqliteWindowFunctionConfig": AiosqliteWindowFunctionConfig, "Literal": Literal, "PathLike": PathLike, }) @@ -467,6 +483,9 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) + _validate_entries( + runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" + ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/aiosqlite/core.py b/sqlspec/adapters/aiosqlite/core.py index 7f449e7fa..201b8a883 100644 --- a/sqlspec/adapters/aiosqlite/core.py +++ b/sqlspec/adapters/aiosqlite/core.py @@ -43,6 +43,7 @@ _T = TypeVar("_T") __all__ = ( + "SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", "AiosqliteStreamSource", "apply_driver_features", "apply_extension_pragmas", @@ -54,6 +55,7 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "end_transaction", "execute_and_resolve_metadata", "execute_and_resolve_rowcount", "execute_fetchall_with_description", @@ -98,12 +100,77 @@ SQLITE_INTERRUPT_CODE = 9 SQLITE_PERM_CODE = 3 SQLITE_READONLY_CODE = 8 +SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT = sys.version_info >= (3, 12) SQLITE_DATABASE_LIST_MIN_COLUMNS = 2 SQLITE_TABLE_LIST_MIN_COLUMNS = 5 SQLITE_TABLE_INFO_MIN_COLUMNS = 2 SQLITE_ROWID_ALIASES = ("rowid", "_rowid_", "oid") +def _end_transaction_on_worker(raw_conn: Any, *, commit: bool, supports_autocommit: bool) -> None: + """End an open transaction on a sqlite3 connection on the worker thread.""" + if supports_autocommit and getattr(raw_conn, "autocommit", False) is True: + if getattr(raw_conn, "in_transaction", False): + cursor = raw_conn.execute("COMMIT" if commit else "ROLLBACK") + with contextlib.suppress(Exception): + cursor.close() + return + if commit: + raw_conn.commit() + else: + raw_conn.rollback() + + +def _read_transaction_state_on_worker(raw_conn: Any) -> "tuple[bool, bool]": + """Read sqlite3 autocommit and in_transaction state on the worker thread.""" + return getattr(raw_conn, "autocommit", False) is True, bool(getattr(raw_conn, "in_transaction", False)) + + +async def end_transaction( + connection: "AiosqliteConnection | Any", *, commit: bool, supports_autocommit: "bool | None" = None +) -> None: + """End an open transaction on an aiosqlite connection. + + Connection.commit and Connection.rollback are no-ops while the underlying + sqlite3 connection runs in autocommit mode, so the statement is issued + directly there. + + Args: + connection: Connection whose transaction should end. + commit: Whether to commit rather than roll back. + supports_autocommit: Whether this runtime's sqlite3 exposes autocommit. + """ + if supports_autocommit is None: + supports_autocommit = SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT + raw_conn = getattr(connection, "_conn", None) + if isinstance(raw_conn, sqlite3.Connection) and callable(getattr(connection, "_execute", None)): + conn_dict = getattr(connection, "__dict__", None) + has_method_override = isinstance(conn_dict, dict) and any( + key in conn_dict for key in ("commit", "rollback", "execute") + ) + if not has_method_override: + await run_on_worker_thread( + connection, _end_transaction_on_worker, raw_conn, commit=commit, supports_autocommit=supports_autocommit + ) + return + autocommit, in_transaction = await run_on_worker_thread(connection, _read_transaction_state_on_worker, raw_conn) + else: + autocommit = getattr(raw_conn, "autocommit", False) is True or getattr(connection, "autocommit", False) is True + in_transaction = ( + bool(getattr(raw_conn, "in_transaction", False)) + if raw_conn is not None + else bool(getattr(connection, "in_transaction", False)) + ) + if supports_autocommit and autocommit: + if in_transaction: + await connection.execute("COMMIT" if commit else "ROLLBACK") + return + if commit: + await connection.commit() + else: + await connection.rollback() + + async def run_on_worker_thread( connection: "AiosqliteConnection", function: "Callable[..., _T]", *args: Any, **kwargs: Any ) -> _T: diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index cfca451ba..775a72e6e 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -4,7 +4,7 @@ import contextlib import inspect import secrets -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Literal, cast from sqlspec.adapters.aiosqlite._typing import AiosqliteCursor, AiosqliteRawCursor, AiosqliteSessionContext from sqlspec.adapters.aiosqlite._typing import aiosqlite_module as aiosqlite @@ -16,6 +16,7 @@ create_mapped_exception, default_statement_config, driver_profile, + end_transaction, execute_and_resolve_metadata, execute_fetchall_with_metadata, execute_many_on_worker_thread, @@ -216,18 +217,26 @@ async def execute_many( await close_result return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - async def begin(self) -> None: - """Begin a database transaction.""" + async def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: + """Begin a database transaction. + + Args: + mode: Transaction locking mode (DEFERRED, IMMEDIATE, EXCLUSIVE). + Falls back to ``driver_features['default_transaction_mode']``, + or ``IMMEDIATE`` when neither is set. + """ + transaction_mode = mode or self.driver_features.get("default_transaction_mode") or "IMMEDIATE" + stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" try: if not self.connection.in_transaction: - await self.connection.execute("BEGIN IMMEDIATE") + await self.connection.execute(stmt) except aiosqlite.Error as e: - await _retry_begin_with_backoff(self.connection, e) + await _retry_begin_with_backoff(self.connection, e, statement=stmt) async def commit(self) -> None: """Commit the current transaction.""" try: - await self.connection.commit() + await end_transaction(self.connection, commit=True) except aiosqlite.Error as e: msg = f"Failed to commit transaction: {e}" raise SQLSpecError(msg) from e @@ -235,7 +244,7 @@ async def commit(self) -> None: async def rollback(self) -> None: """Rollback the current transaction.""" try: - await self.connection.rollback() + await end_transaction(self.connection, commit=False) except aiosqlite.Error as e: msg = f"Failed to rollback transaction: {e}" raise SQLSpecError(msg) from e @@ -545,9 +554,13 @@ def _connection_in_transaction(self) -> bool: async def _retry_begin_with_backoff( - connection: "AiosqliteConnection", initial_error: aiosqlite.Error, max_retries: int = 3 + connection: "AiosqliteConnection", + initial_error: aiosqlite.Error, + max_retries: int = 3, + *, + statement: str = "BEGIN IMMEDIATE", ) -> None: - """Retry ``BEGIN IMMEDIATE`` after SQLite reports a busy connection. + """Retry transaction start 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 @@ -557,6 +570,7 @@ async def _retry_begin_with_backoff( 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. + statement: SQL statement used to start the transaction. Raises: SQLSpecError: If every retry attempt fails. @@ -565,7 +579,7 @@ async def _retry_begin_with_backoff( delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) await asyncio.sleep(delay) try: - await connection.execute("BEGIN IMMEDIATE") + await connection.execute(statement) except aiosqlite.Error: if attempt == max_retries - 1: break diff --git a/sqlspec/adapters/aiosqlite/litestar/store.py b/sqlspec/adapters/aiosqlite/litestar/store.py index ddffcafba..d53de75bf 100644 --- a/sqlspec/adapters/aiosqlite/litestar/store.py +++ b/sqlspec/adapters/aiosqlite/litestar/store.py @@ -5,7 +5,7 @@ from typing_extensions import NotRequired -from sqlspec.adapters.aiosqlite.core import apply_extension_pragmas, extension_pragma_statements +from sqlspec.adapters.aiosqlite.core import apply_extension_pragmas, end_transaction, extension_pragma_statements from sqlspec.config import LitestarConfig from sqlspec.extensions.litestar.store import BaseSQLSpecStore @@ -112,7 +112,7 @@ async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "by WHERE session_id = ? """ await conn.execute(update_sql, (new_expires_at_julian, key)) - await conn.commit() + await end_transaction(conn, commit=True) return bytes(data) @@ -135,7 +135,7 @@ async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta async with self._config.provide_connection() as conn: await conn.execute(sql, (key, data, expires_at_julian)) - await conn.commit() + await end_transaction(conn, commit=True) async def delete(self, key: str) -> None: """Delete a session by key. @@ -147,7 +147,7 @@ async def delete(self, key: str) -> None: async with self._config.provide_connection() as conn: await conn.execute(sql, (key,)) - await conn.commit() + await end_transaction(conn, commit=True) async def delete_all(self) -> None: """Delete all sessions from the store.""" @@ -155,7 +155,7 @@ async def delete_all(self) -> None: async with self._config.provide_connection() as conn: await conn.execute(sql) - await conn.commit() + await end_transaction(conn, commit=True) self._log_delete_all() async def exists(self, key: str) -> bool: @@ -218,12 +218,15 @@ async def delete_expired(self) -> int: Returns: Number of sessions deleted. """ - sql = f"DELETE FROM {self._table_name} WHERE expires_at <= julianday('now')" + sql = f"DELETE FROM {self._table_name} WHERE expires_at IS NOT NULL AND expires_at <= julianday('now')" async with self._config.provide_connection() as conn: cursor = await conn.execute(sql) - await conn.commit() - count = cursor.rowcount + try: + await end_transaction(conn, commit=True) + count = cursor.rowcount + finally: + await cursor.close() if count > 0: self._log_delete_expired(count) return count diff --git a/sqlspec/adapters/aiosqlite/pool.py b/sqlspec/adapters/aiosqlite/pool.py index 23e577bb9..c1a0066ae 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -10,7 +10,7 @@ from sqlspec.adapters.aiosqlite._typing import aiosqlite_module as aiosqlite from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 -from sqlspec.adapters.aiosqlite.core import run_on_worker_thread +from sqlspec.adapters.aiosqlite.core import end_transaction, run_on_worker_thread from sqlspec.exceptions import SQLSpecError from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -78,9 +78,12 @@ def _has_active_transaction(connection: "AiosqliteConnection") -> bool: def _register_runtime_objects( - connection: "AiosqliteConnection", aggregates: "Sequence[dict[str, Any]]", collations: "Sequence[dict[str, Any]]" + connection: "AiosqliteConnection", + aggregates: "Sequence[dict[str, Any]]", + collations: "Sequence[dict[str, Any]]", + window_functions: "Sequence[dict[str, Any]]" = (), ) -> None: - """Register custom aggregates and collations on the worker thread.""" + """Register custom aggregates, collations, and window functions on the worker thread.""" raw_connection = connection._conn for aggregate_config in aggregates: raw_connection.create_aggregate( @@ -88,6 +91,10 @@ def _register_runtime_objects( ) for collation_config in collations: raw_connection.create_collation(collation_config["name"], collation_config["func"]) + create_window_fn = getattr(raw_connection, "create_window_function", None) + if create_window_fn is not None: + for window_config in window_functions: + create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: "dict[str, Any]") -> None: @@ -115,8 +122,11 @@ async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: aggregates = runtime_setup.get("custom_aggregates", ()) collations = runtime_setup.get("custom_collations", ()) - if aggregates or collations: - await run_on_worker_thread(connection, _register_runtime_objects, connection, aggregates, collations) + window_functions = runtime_setup.get("custom_window_functions", ()) + if aggregates or collations or window_functions: + await run_on_worker_thread( + connection, _register_runtime_objects, connection, aggregates, collations, window_functions + ) authorizer_callback = runtime_setup.get("authorizer_callback") if authorizer_callback is not None: @@ -228,7 +238,7 @@ async def reset(self) -> None: if not _has_active_transaction(self.connection): return with suppress(Exception): - await self.connection.rollback() + await end_transaction(self.connection, commit=False) async def close(self) -> None: """Close the connection.""" @@ -237,7 +247,7 @@ async def close(self) -> None: try: if _has_active_transaction(self.connection): with suppress(Exception): - await self.connection.rollback() + await end_transaction(self.connection, commit=False) await self.connection.close() except Exception: log_with_context( @@ -805,7 +815,7 @@ async def release(self, connection: AiosqlitePoolConnection) -> None: try: if _has_active_transaction(connection.connection): with suppress(Exception): - await connection.connection.rollback() + await end_transaction(connection.connection, commit=False) connection.idle_since = time.time() self._queue.put_nowait(connection) except Exception as e: diff --git a/sqlspec/adapters/sqlite/__init__.py b/sqlspec/adapters/sqlite/__init__.py index 4a17a6019..f534e9894 100644 --- a/sqlspec/adapters/sqlite/__init__.py +++ b/sqlspec/adapters/sqlite/__init__.py @@ -8,6 +8,7 @@ SqliteConnectionParams, SqliteDriverFeatures, SqliteFunctionConfig, + SqliteWindowFunctionConfig, ) from sqlspec.adapters.sqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.sqlite.driver import SqliteDriver, SqliteExceptionHandler @@ -25,6 +26,7 @@ "SqliteDriverFeatures", "SqliteExceptionHandler", "SqliteFunctionConfig", + "SqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index 75b26ceaf..5e0b17f25 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -1,6 +1,7 @@ """SQLite driver implementation.""" -from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast +import contextlib +from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast from mypy_extensions import mypyc_attr @@ -54,7 +55,6 @@ __all__ = ("SqliteCursor", "SqliteDriver", "SqliteExceptionHandler", "SqliteSessionContext") T = TypeVar("T") -_BATCH_SAMPLE_THRESHOLD: Final = 100 @mypyc_attr(allow_interpreted_subclasses=True) @@ -221,13 +221,17 @@ def execute_many( and self.observability.is_idle and self._can_use_execute_many_thin_path(statement, parameters, config) ): + cursor = None try: cursor = self.connection.executemany(statement, parameters) + affected_rows = resolve_rowcount(cursor) except sqlite3.Error as exc: raise create_mapped_exception(exc) from exc + finally: + if cursor is not None: + with contextlib.suppress(Exception): + cursor.close() - 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) @@ -432,6 +436,7 @@ def _execute_cache_hit( This bypasses cursor context-manager overhead for repeated cached statements while preserving driver exception mapping behavior. """ + cursor = None direct_statement: SQL | None = None returns_rows = cached.operation_profile.returns_rows self._invalidate_rowid_target_cache(cached.operation_type) @@ -483,6 +488,9 @@ def _execute_cache_hit( ) return self.build_statement_result(direct_statement, execution_result) finally: + if cursor is not None: + with contextlib.suppress(Exception): + cursor.close() if direct_statement is not None: self._release_pooled_statement(direct_statement) msg = "unreachable" @@ -521,7 +529,7 @@ def _can_use_execute_many_thin_path( @staticmethod def _thin_path_parameters_are_eligible( - parameters: "list[StatementParameters]", type_coercion_map: "dict[type, Any] | None" + parameters: "Sequence[StatementParameters]", type_coercion_map: "dict[type, Any] | None" ) -> bool: """Validate parameter payload for the SQLite execute-many thin path.""" first_sequence = SqliteDriver._as_sequence_parameter_set(parameters[0]) @@ -534,16 +542,9 @@ def _thin_path_parameters_are_eligible( has_type_coercion = bool(coercion_map) fallback_items = type_coercion_fallbacks(coercion_map) if coercion_map else () - total_rows = len(parameters) - if total_rows > _BATCH_SAMPLE_THRESHOLD: - sample_indices = (0, 1, total_rows // 4, total_rows // 2, (3 * total_rows) // 4, total_rows - 1) - eval_parameters = [parameters[i] for i in sample_indices] - else: - eval_parameters = parameters - if row_len == 1: if has_type_coercion and coercion_map is not None: - for param_set in eval_parameters: + for param_set in parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -553,7 +554,7 @@ def _thin_path_parameters_are_eligible( return False return True - for param_set in eval_parameters: + for param_set in parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -564,7 +565,7 @@ def _thin_path_parameters_are_eligible( return True if has_type_coercion and coercion_map is not None: - for param_set in eval_parameters: + for param_set in parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False @@ -575,7 +576,7 @@ def _thin_path_parameters_are_eligible( return False return True - for param_set in eval_parameters: + for param_set in parameters: sequence = SqliteDriver._as_sequence_parameter_set(param_set) if sequence is None or type(sequence) is not first_type: return False diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 16b5b291d..540f43c86 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -106,6 +106,8 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", + "custom_window_functions", + "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", diff --git a/tests/unit/adapters/test_aiosqlite/test_config.py b/tests/unit/adapters/test_aiosqlite/test_config.py index 4b13f2eb7..a17c48a2a 100644 --- a/tests/unit/adapters/test_aiosqlite/test_config.py +++ b/tests/unit/adapters/test_aiosqlite/test_config.py @@ -237,6 +237,11 @@ def test_custom_aggregate_entry_requires_keys() -> None: AiosqliteConfig(driver_features={"custom_aggregates": [{"name": "agg", "narg": 1}]}) +def test_custom_window_function_entry_requires_keys() -> None: + with pytest.raises(ImproperConfigurationError, match="custom_window_functions"): + AiosqliteConfig(driver_features={"custom_window_functions": [{"name": "win", "narg": 1}]}) + + def test_progress_handler_interval_must_be_positive() -> None: with pytest.raises(ImproperConfigurationError, match="progress_handler_interval"): AiosqliteConfig(driver_features={"progress_handler": lambda: None, "progress_handler_interval": 0}) diff --git a/tests/unit/adapters/test_aiosqlite/test_driver.py b/tests/unit/adapters/test_aiosqlite/test_driver.py index 9f794349b..27fff531c 100644 --- a/tests/unit/adapters/test_aiosqlite/test_driver.py +++ b/tests/unit/adapters/test_aiosqlite/test_driver.py @@ -345,3 +345,74 @@ def test_profile_aiosqlite_statement_config_parity_with_sqlite() -> None: aio_config = build_statement_config() sqlite_config = sqlite_build_statement_config() assert aio_config.enable_parameter_type_wrapping == sqlite_config.enable_parameter_type_wrapping + + +class _AsyncAutocommitConnection: + """Mimics an aiosqlite connection wrapping a Python 3.12+ sqlite3 connection in autocommit mode.""" + + def __init__(self, in_transaction: bool = True, autocommit: bool = True) -> None: + self.autocommit = autocommit + self.in_transaction = in_transaction + self._conn = self + self.statements: list[str] = [] + self.commit_calls = 0 + self.rollback_calls = 0 + + async def execute(self, sql: str, parameters: object = ()) -> None: + _ = parameters + self.statements.append(sql) + + async def commit(self) -> None: + self.commit_calls += 1 + + async def rollback(self) -> None: + self.rollback_calls += 1 + + +@pytest.mark.parametrize( + ("method", "statement"), [("commit", "COMMIT"), ("rollback", "ROLLBACK")], ids=["commit", "rollback"] +) +async def test_aiosqlite_autocommit_mode_ends_transactions_with_explicit_statement( + method: str, statement: str, monkeypatch: pytest.MonkeyPatch +) -> None: + """In autocommit mode AiosqliteDriver commit/rollback must execute explicit SQL.""" + monkeypatch.setattr("sqlspec.adapters.aiosqlite.core.SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", True) + connection = _AsyncAutocommitConnection(in_transaction=True, autocommit=True) + driver = AiosqliteDriver(connection=cast("Any", connection), statement_config=default_statement_config) + + await getattr(driver, method)() + + assert connection.statements == [statement] + assert connection.commit_calls == 0 + assert connection.rollback_calls == 0 + + +@pytest.mark.parametrize("method", ["commit", "rollback"]) +async def test_aiosqlite_autocommit_mode_skips_statement_without_open_transaction( + method: str, monkeypatch: pytest.MonkeyPatch +) -> None: + """Ending a transaction in autocommit mode without an open transaction must be a no-op.""" + monkeypatch.setattr("sqlspec.adapters.aiosqlite.core.SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", True) + connection = _AsyncAutocommitConnection(in_transaction=False, autocommit=True) + driver = AiosqliteDriver(connection=cast("Any", connection), statement_config=default_statement_config) + + await getattr(driver, method)() + + assert connection.statements == [] + assert connection.commit_calls == 0 + assert connection.rollback_calls == 0 + + +async def test_aiosqlite_begin_honors_explicit_and_default_transaction_mode() -> None: + """AiosqliteDriver.begin should honor explicit mode and driver_features default_transaction_mode.""" + connection = _AsyncAutocommitConnection(in_transaction=False, autocommit=False) + driver = AiosqliteDriver( + connection=cast("Any", connection), + statement_config=default_statement_config, + driver_features={"default_transaction_mode": "DEFERRED"}, + ) + + await driver.begin() + await driver.begin(mode="EXCLUSIVE") + + assert connection.statements == ["BEGIN DEFERRED", "BEGIN EXCLUSIVE"] diff --git a/tests/unit/adapters/test_sqlite/test_driver.py b/tests/unit/adapters/test_sqlite/test_driver.py index d8b37891e..5f0098696 100644 --- a/tests/unit/adapters/test_sqlite/test_driver.py +++ b/tests/unit/adapters/test_sqlite/test_driver.py @@ -1,11 +1,12 @@ import sqlite3 +from collections import defaultdict from pathlib import Path from typing import Any, cast import pytest from sqlspec.adapters.sqlite import SqliteDriver -from sqlspec.adapters.sqlite.core import SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT +from sqlspec.adapters.sqlite.core import SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT, default_statement_config def test_rowid_eligibility_falls_back_when_table_list_is_unavailable() -> None: @@ -179,3 +180,76 @@ def test_legacy_mode_still_uses_the_connection_methods( assert connection.statements == [] assert getattr(connection, attribute) == 1 + + +class _TrackingCursor: + def __init__(self, rows: list[tuple[Any, ...]] | None = None) -> None: + self._rows = rows or [] + self.description = [("id",)] if rows is not None else None + self.rowcount = len(self._rows) if self._rows else 2 + self.lastrowid = 1 + self.closed = False + + def fetchall(self) -> list[tuple[Any, ...]]: + return self._rows + + def close(self) -> None: + self.closed = True + + +class _TrackingConnection: + def __init__(self) -> None: + self.in_transaction = False + self.cursors: list[_TrackingCursor] = [] + + def execute(self, sql: str, parameters: object = ()) -> _TrackingCursor: + _ = (sql, parameters) + cursor = _TrackingCursor(rows=[(1,)]) + self.cursors.append(cursor) + return cursor + + def executemany(self, sql: str, parameters: object) -> _TrackingCursor: + _ = (sql, parameters) + cursor = _TrackingCursor() + self.cursors.append(cursor) + return cursor + + +def test_execute_many_thin_path_closes_cursor() -> None: + """SqliteDriver.execute_many thin path must close its cursor in finally.""" + connection = _TrackingConnection() + driver = SqliteDriver(connection=cast("Any", connection)) + + result = driver.execute_many("INSERT INTO items (name) VALUES (?)", [("a",), ("b",)]) + + assert result.rows_affected == 2 + assert len(connection.cursors) == 1 + assert connection.cursors[0].closed is True + + +def test_execute_cache_hit_closes_cursor() -> None: + """SqliteDriver._execute_cache_hit must close its cursor on cached execution.""" + connection = sqlite3.connect(":memory:") + connection.execute("CREATE TABLE items (id INTEGER PRIMARY KEY, name TEXT)") + connection.execute("INSERT INTO items (name) VALUES ('alpha')") + driver = SqliteDriver(connection=connection) + try: + first = driver.execute("SELECT id, name FROM items WHERE id = ?", (1,)) + second = driver.execute("SELECT id, name FROM items WHERE id = ?", (1,)) + assert first.get_data() == [{"id": 1, "name": "alpha"}] + assert second.get_data() == [{"id": 1, "name": "alpha"}] + finally: + connection.close() + + +def test_execute_many_thin_path_checks_all_rows_beyond_sample_threshold() -> None: + """_thin_path_parameters_are_eligible must inspect every row even in large batches.""" + rows: list[tuple[Any, ...]] = [(i,) for i in range(120)] + rows[57] = (defaultdict(int, a=1),) + + assert ( + SqliteDriver._thin_path_parameters_are_eligible( + rows, default_statement_config.parameter_config.type_coercion_map + ) + is False + ) From 14693b709daa205503a71c7c529d6aa2e039dca9 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 20:53:34 +0000 Subject: [PATCH 08/12] fix(sqlite): preserve adapter contracts during optimization review --- pyproject.toml | 1 - sqlspec/adapters/aiosqlite/__init__.py | 2 - sqlspec/adapters/aiosqlite/_typing.py | 22 +++++-- sqlspec/adapters/aiosqlite/config.py | 19 ------ sqlspec/adapters/aiosqlite/core.py | 35 ++--------- sqlspec/adapters/aiosqlite/driver.py | 61 ++++++++----------- sqlspec/adapters/aiosqlite/pool.py | 30 +++------ sqlspec/adapters/aiosqlite/type_converter.py | 20 +++--- sqlspec/adapters/sqlite/__init__.py | 2 - sqlspec/adapters/sqlite/_typing.py | 21 ++++--- sqlspec/adapters/sqlite/config.py | 19 ------ sqlspec/adapters/sqlite/core.py | 9 +-- sqlspec/adapters/sqlite/driver.py | 31 ++++------ sqlspec/adapters/sqlite/pool.py | 11 +--- sqlspec/adapters/sqlite/type_converter.py | 20 +++--- .../adapters/test_aiosqlite/test_driver.py | 15 ----- .../unit/adapters/test_aiosqlite/test_pool.py | 14 +++++ .../unit/adapters/test_sqlite/test_driver.py | 30 +++++++++ tests/unit/adapters/test_sqlite/test_pool.py | 11 ++++ tools/scripts/mypyc_inventory.py | 4 +- 20 files changed, 160 insertions(+), 217 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index d9b4b0e80..6158692aa 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -292,7 +292,6 @@ include = [ "sqlspec/adapters/mysql_common.py", # Shared MySQL-family adapter helpers "sqlspec/adapters/**/core.py", # Adapter compiled helpers "sqlspec/adapters/**/type_converter.py", # All adapters type converters - "sqlspec/adapters/sqlite/driver.py", # SQLite synchronous driver "sqlspec/adapters/oracledb/_param_types.py", # Slot-based LOB/JSON parameter wrappers "sqlspec/adapters/oracledb/_json_handlers.py", # Native JSON inputtypehandler / outputtypehandler chain "sqlspec/adapters/oracledb/_uuid_handlers.py", # UUID ↔ RAW(16) inputtypehandler / outputtypehandler chain diff --git a/sqlspec/adapters/aiosqlite/__init__.py b/sqlspec/adapters/aiosqlite/__init__.py index daa39414d..1f48983b9 100644 --- a/sqlspec/adapters/aiosqlite/__init__.py +++ b/sqlspec/adapters/aiosqlite/__init__.py @@ -7,7 +7,6 @@ AiosqliteDriverFeatures, AiosqliteFunctionConfig, AiosqlitePoolParams, - AiosqliteWindowFunctionConfig, ) from sqlspec.adapters.aiosqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver, AiosqliteExceptionHandler @@ -35,7 +34,6 @@ "AiosqlitePoolConnection", "AiosqlitePoolParams", "AiosqliteRawCursor", - "AiosqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/aiosqlite/_typing.py b/sqlspec/adapters/aiosqlite/_typing.py index e908351e6..55b6f9588 100644 --- a/sqlspec/adapters/aiosqlite/_typing.py +++ b/sqlspec/adapters/aiosqlite/_typing.py @@ -1,3 +1,4 @@ +# pyright: reportCallIssue=false, reportAttributeAccessIssue=false, reportArgumentType=false """AIOSQLite adapter type definitions. This module contains type aliases and classes that are excluded from mypyc @@ -5,28 +6,37 @@ """ import contextlib +import sqlite3 import sqlite3 as aiosqlite_sqlite_module -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, Any +import aiosqlite import aiosqlite as aiosqlite_module -from aiosqlite import Connection as AiosqliteConnection -from aiosqlite import Cursor as AiosqliteRawCursor -from aiosqlite import Error as AiosqliteError +from typing_extensions import TypeAliasType -AiosqliteConnectionFactory: TypeAlias = type[aiosqlite_sqlite_module.Connection] +_AiosqliteConnection = aiosqlite.Connection if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver from sqlspec.core import StatementConfig + AiosqliteConnection: TypeAlias = _AiosqliteConnection + AiosqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] + AiosqliteRawCursor: TypeAlias = aiosqlite.Cursor + +if not TYPE_CHECKING: + AiosqliteConnection = _AiosqliteConnection + AiosqliteConnectionFactory = TypeAliasType("AiosqliteConnectionFactory", type[sqlite3.Connection]) + AiosqliteRawCursor = aiosqlite.Cursor + __all__ = ( "AiosqliteConnection", "AiosqliteConnectionFactory", "AiosqliteCursor", - "AiosqliteError", "AiosqliteRawCursor", "AiosqliteSessionContext", "aiosqlite_module", diff --git a/sqlspec/adapters/aiosqlite/config.py b/sqlspec/adapters/aiosqlite/config.py index d5c80338d..5741c9e7b 100644 --- a/sqlspec/adapters/aiosqlite/config.py +++ b/sqlspec/adapters/aiosqlite/config.py @@ -48,7 +48,6 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", - "AiosqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -110,14 +109,6 @@ class AiosqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" -class AiosqliteWindowFunctionConfig(TypedDict): - """User-defined aiosqlite window function registration.""" - - name: str - narg: int - window_class: "type[Any]" - - class AiosqliteDriverFeatures(TypedDict): """Aiosqlite driver feature configuration. @@ -147,9 +138,6 @@ class AiosqliteDriverFeatures(TypedDict): Each entry must include name and func. Callable values are plain sync callables. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. - custom_window_functions: Register user-defined aggregate window functions. - Each entry must include name, narg, and window_class. - default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation on the worker thread. trace_callback: sqlite3 trace hook run for executed statements on the worker thread. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -173,8 +161,6 @@ class AiosqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[AiosqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[AiosqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[AiosqliteAggregateConfig]]" - custom_window_functions: "NotRequired[Sequence[AiosqliteWindowFunctionConfig]]" - default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -191,7 +177,6 @@ class AiosqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", - "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -363,7 +348,6 @@ def get_signature_namespace(self) -> "dict[str, Any]": "AiosqliteFunctionConfig": AiosqliteFunctionConfig, "AiosqlitePoolParams": AiosqlitePoolParams, "AiosqliteSessionContext": AiosqliteSessionContext, - "AiosqliteWindowFunctionConfig": AiosqliteWindowFunctionConfig, "Literal": Literal, "PathLike": PathLike, }) @@ -483,9 +467,6 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) - _validate_entries( - runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" - ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/aiosqlite/core.py b/sqlspec/adapters/aiosqlite/core.py index 201b8a883..ecb8a7e75 100644 --- a/sqlspec/adapters/aiosqlite/core.py +++ b/sqlspec/adapters/aiosqlite/core.py @@ -61,7 +61,6 @@ "execute_fetchall_with_description", "execute_fetchall_with_metadata", "execute_many_on_worker_thread", - "execute_script_on_worker_thread", "extension_pragma_statements", "format_identifier", "normalize_execute_many_parameters", @@ -325,16 +324,15 @@ def normalize_execute_parameters(parameters: Any) -> Any: class AiosqliteStreamSource: - """Compiled async chunk source streaming dict or tuple rows from an aiosqlite cursor via ``fetchmany``.""" + """Compiled async chunk source streaming dict rows from an aiosqlite cursor via ``fetchmany``.""" - __slots__ = ("_as_dict", "_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") + __slots__ = ("_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") - def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int, as_dict: bool = True) -> None: + def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: self._driver = driver self._sql = sql self._parameters = parameters self._chunk_size = chunk_size - self._as_dict = as_dict self._cursor: Any = None self._column_names: list[str] | None = None @@ -349,16 +347,12 @@ async def _start(self) -> None: self._cursor = cursor await cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) - async def fetch_chunk(self) -> "list[Any]": + async def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() - rows: list[Any] = await self._driver._run_with_exception_handler( - handler, self._cursor.fetchmany, self._chunk_size - ) + rows = await self._driver._run_with_exception_handler(handler, self._cursor.fetchmany, self._chunk_size) self._driver._check_pending_exception(handler) if not rows: return [] - if not self._as_dict: - return rows if self._column_names is None: self._column_names = [description[0] for description in self._cursor.description] return rows_to_dicts(rows, self._column_names) @@ -672,28 +666,9 @@ def _execute_many_on_worker_thread(connection: "AiosqliteConnection", sql: str, cast("Any", cursor).close() -def _execute_script_on_worker_thread( - connection: "AiosqliteConnection", statements: "Sequence[str]", parameters: Any -) -> tuple[int, int]: - """Execute multi-statement SQL script on the worker thread.""" - raw_connection = connection._conn - cursor = raw_connection.cursor() - normalized_params = normalize_execute_parameters(parameters) - successful_count = 0 - try: - for stmt in statements: - cursor.execute(stmt, normalized_params) - successful_count += 1 - return len(statements), successful_count - finally: - with contextlib.suppress(Exception): - cursor.close() - - execute_and_resolve_metadata = _execute_and_resolve_metadata execute_fetchall_with_metadata = _execute_fetchall_with_metadata execute_many_on_worker_thread = _execute_many_on_worker_thread -execute_script_on_worker_thread = _execute_script_on_worker_thread def _resolve_insert_target(expression: Any) -> "tuple[str | None, str] | None": diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index 775a72e6e..e506fed64 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -3,8 +3,8 @@ import asyncio import contextlib import inspect -import secrets -from typing import TYPE_CHECKING, Any, Literal, cast +import random +from typing import TYPE_CHECKING, Any, cast from sqlspec.adapters.aiosqlite._typing import AiosqliteCursor, AiosqliteRawCursor, AiosqliteSessionContext from sqlspec.adapters.aiosqlite._typing import aiosqlite_module as aiosqlite @@ -20,7 +20,6 @@ execute_and_resolve_metadata, execute_fetchall_with_metadata, execute_many_on_worker_thread, - execute_script_on_worker_thread, format_identifier, normalize_execute_parameters, resolve_rowcount, @@ -58,7 +57,6 @@ "AiosqliteSessionContext", ) -random = secrets.SystemRandom() _execute_and_resolve_metadata = execute_and_resolve_metadata _execute_fetchall_with_metadata = execute_fetchall_with_metadata @@ -169,15 +167,18 @@ async def dispatch_execute_script(self, cursor: "AiosqliteRawCursor", statement: 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: - statement_count, successful_count = await run_on_worker_thread( - self.connection, execute_script_on_worker_thread, self.connection, statements, prepared_parameters - ) + 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( - cursor, statement_count=statement_count, successful_statements=successful_count, is_script_result=True + last_cursor, statement_count=len(statements), successful_statements=successful_count, is_script_result=True ) async def execute_many( @@ -217,21 +218,13 @@ async def execute_many( await close_result return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - async def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: - """Begin a database transaction. - - Args: - mode: Transaction locking mode (DEFERRED, IMMEDIATE, EXCLUSIVE). - Falls back to ``driver_features['default_transaction_mode']``, - or ``IMMEDIATE`` when neither is set. - """ - transaction_mode = mode or self.driver_features.get("default_transaction_mode") or "IMMEDIATE" - stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" + async def begin(self) -> None: + """Begin a database transaction.""" try: if not self.connection.in_transaction: - await self.connection.execute(stmt) + await self.connection.execute("BEGIN IMMEDIATE") except aiosqlite.Error as e: - await _retry_begin_with_backoff(self.connection, e, statement=stmt) + await _retry_begin_with_backoff(self.connection, e) async def commit(self) -> None: """Commit the current transaction.""" @@ -253,14 +246,12 @@ def with_cursor(self, connection: "AiosqliteConnection") -> "AiosqliteCursor": """Create async context manager for AIOSQLite cursor.""" return AiosqliteCursor(connection) - def dispatch_select_stream( - self, statement: "SQL", chunk_size: int, as_dict: bool = True - ) -> "AsyncRowStream[Any] | None": + 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, as_dict=as_dict)) + return AsyncRowStream(AiosqliteStreamSource(self, sql, prepared_parameters, chunk_size)) def handle_database_exceptions(self) -> "AiosqliteExceptionHandler": """Handle AIOSQLite-specific exceptions.""" @@ -294,7 +285,6 @@ async def load_from_arrow( table: str, source: "ArrowResult | Any", *, - batch_size: int = 10000, partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, @@ -314,7 +304,7 @@ async def load_from_arrow( statement = f"DELETE FROM {format_identifier(table)}" async with self.with_cursor(self.connection) as cursor: await cursor.execute(statement) - for batch in arrow_table.to_batches(max_chunksize=batch_size): + for batch in arrow_table.to_batches(max_chunksize=10000): pydict = batch.to_pydict() records = list(zip(*(pydict[col] for col in columns), strict=False)) if records: @@ -327,10 +317,12 @@ async def load_from_arrow( await cursor.executemany(insert_sql, cast("Any", prepared_records)) if owns_transaction: await self.commit() - except (aiosqlite.Error, sqlite3.Error) as exc: + except BaseException as exc: if owns_transaction: await self.rollback() - raise create_mapped_exception(exc) from exc + if isinstance(exc, (aiosqlite.Error, sqlite3.Error)): + raise create_mapped_exception(exc) from exc + raise telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table @@ -554,13 +546,9 @@ def _connection_in_transaction(self) -> bool: async def _retry_begin_with_backoff( - connection: "AiosqliteConnection", - initial_error: aiosqlite.Error, - max_retries: int = 3, - *, - statement: str = "BEGIN IMMEDIATE", + connection: "AiosqliteConnection", initial_error: aiosqlite.Error, max_retries: int = 3 ) -> None: - """Retry transaction start after SQLite reports a busy connection. + """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 @@ -570,16 +558,15 @@ async def _retry_begin_with_backoff( 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. - statement: SQL statement used to start the transaction. Raises: SQLSpecError: If every retry attempt fails. """ for attempt in range(max_retries): - delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) + delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) # noqa: S311 await asyncio.sleep(delay) try: - await connection.execute(statement) + await connection.execute("BEGIN IMMEDIATE") except aiosqlite.Error: if attempt == max_retries - 1: break diff --git a/sqlspec/adapters/aiosqlite/pool.py b/sqlspec/adapters/aiosqlite/pool.py index c1a0066ae..67f93bbd2 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -78,12 +78,9 @@ def _has_active_transaction(connection: "AiosqliteConnection") -> bool: def _register_runtime_objects( - connection: "AiosqliteConnection", - aggregates: "Sequence[dict[str, Any]]", - collations: "Sequence[dict[str, Any]]", - window_functions: "Sequence[dict[str, Any]]" = (), + connection: "AiosqliteConnection", aggregates: "Sequence[dict[str, Any]]", collations: "Sequence[dict[str, Any]]" ) -> None: - """Register custom aggregates, collations, and window functions on the worker thread.""" + """Register custom aggregates and collations on the worker thread.""" raw_connection = connection._conn for aggregate_config in aggregates: raw_connection.create_aggregate( @@ -91,10 +88,6 @@ def _register_runtime_objects( ) for collation_config in collations: raw_connection.create_collation(collation_config["name"], collation_config["func"]) - create_window_fn = getattr(raw_connection, "create_window_function", None) - if create_window_fn is not None: - for window_config in window_functions: - create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: "dict[str, Any]") -> None: @@ -122,11 +115,8 @@ async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: aggregates = runtime_setup.get("custom_aggregates", ()) collations = runtime_setup.get("custom_collations", ()) - window_functions = runtime_setup.get("custom_window_functions", ()) - if aggregates or collations or window_functions: - await run_on_worker_thread( - connection, _register_runtime_objects, connection, aggregates, collations, window_functions - ) + if aggregates or collations: + await run_on_worker_thread(connection, _register_runtime_objects, connection, aggregates, collations) authorizer_callback = runtime_setup.get("authorizer_callback") if authorizer_callback is not None: @@ -394,9 +384,9 @@ def _set_connect_proxy_daemon(self, connect_proxy: Any) -> None: async def _force_stop_connection(self, connection: AiosqlitePoolConnection, *, reason: str) -> None: """Force-stop aiosqlite worker thread when graceful close times out.""" - raw_conn = getattr(connection.connection, "_conn", None) - if raw_conn is not None: - with suppress(Exception): + with suppress(Exception): + raw_conn = getattr(connection.connection, "_conn", None) + if raw_conn is not None: raw_conn.interrupt() try: stop_method = getattr(connection.connection, "stop", None) @@ -850,9 +840,9 @@ async def close(self) -> None: if connections: for conn in connections: - raw_conn = getattr(conn.connection, "_conn", None) - if raw_conn is not None: - with suppress(Exception): + with suppress(Exception): + raw_conn = getattr(conn.connection, "_conn", None) + if raw_conn is not None: raw_conn.interrupt() close_tasks = [asyncio.wait_for(conn.close(), timeout=self._operation_timeout) for conn in connections] results = await asyncio.gather(*close_tasks, return_exceptions=True) diff --git a/sqlspec/adapters/aiosqlite/type_converter.py b/sqlspec/adapters/aiosqlite/type_converter.py index 3579d1d7c..d4fa27f7b 100644 --- a/sqlspec/adapters/aiosqlite/type_converter.py +++ b/sqlspec/adapters/aiosqlite/type_converter.py @@ -1,3 +1,4 @@ +# Keep in sync with sqlspec/adapters/sqlite/type_converter.py """SQLite custom type handlers for optional JSON and type conversion support. Provides registration functions for SQLite's adapter/converter system to enable @@ -8,12 +9,12 @@ instead of lambdas for adapter registration. """ +import json from functools import partial from typing import TYPE_CHECKING, Any from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 from sqlspec.utils.logging import get_logger -from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Callable @@ -30,13 +31,14 @@ def json_adapter(value: Any, serializer: "Callable[[Any], str] | None" = None) - Args: value: Python dict or list to serialize. - serializer: Optional JSON serializer callable. Defaults to compiled to_json. + serializer: Optional JSON serializer callable. Defaults to standard json.dumps. Returns: JSON string representation. """ - codec = serializer or to_json - return codec(value) + if serializer is None: + return json.dumps(value, ensure_ascii=False) + return serializer(value) def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = None) -> Any: @@ -44,13 +46,14 @@ def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = N Args: value: UTF-8 encoded JSON bytes from SQLite. - deserializer: Optional JSON deserializer callable. Defaults to compiled from_json. + deserializer: Optional JSON deserializer callable. Defaults to standard json.loads. Returns: Deserialized Python object (dict or list). """ - codec = deserializer or from_json - return codec(value.decode("utf-8")) + if deserializer is None: + return json.loads(value.decode("utf-8")) + return deserializer(value.decode("utf-8")) def register_type_handlers( @@ -61,9 +64,6 @@ def register_type_handlers( This function registers handlers globally for the sqlite3 module. It should be called once during application initialization if custom type handling is needed. - Note that sqlite3.register_adapter is deprecated in Python 3.12+ in favor of - statement-level conversions. - Args: json_serializer: Optional custom JSON serializer. json_deserializer: Optional custom JSON deserializer. diff --git a/sqlspec/adapters/sqlite/__init__.py b/sqlspec/adapters/sqlite/__init__.py index f534e9894..4a17a6019 100644 --- a/sqlspec/adapters/sqlite/__init__.py +++ b/sqlspec/adapters/sqlite/__init__.py @@ -8,7 +8,6 @@ SqliteConnectionParams, SqliteDriverFeatures, SqliteFunctionConfig, - SqliteWindowFunctionConfig, ) from sqlspec.adapters.sqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.sqlite.driver import SqliteDriver, SqliteExceptionHandler @@ -26,7 +25,6 @@ "SqliteDriverFeatures", "SqliteExceptionHandler", "SqliteFunctionConfig", - "SqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/sqlite/_typing.py b/sqlspec/adapters/sqlite/_typing.py index 7a974adad..25dfe5676 100644 --- a/sqlspec/adapters/sqlite/_typing.py +++ b/sqlspec/adapters/sqlite/_typing.py @@ -5,28 +5,33 @@ """ import contextlib +import sqlite3 import sqlite3 as sqlite_module -from sqlite3 import Connection as SqliteConnection -from sqlite3 import Cursor as SqliteRawCursor -from sqlite3 import Error as SqliteError -from sqlite3 import OperationalError as SqliteOperationalError -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, Any -SqliteConnectionFactory: TypeAlias = type[SqliteConnection] +_SqliteConnection = sqlite3.Connection if TYPE_CHECKING: from collections.abc import Callable from types import TracebackType + from typing import TypeAlias from sqlspec.adapters.sqlite.driver import SqliteDriver from sqlspec.core import StatementConfig + SqliteConnection: TypeAlias = _SqliteConnection + SqliteConnectionFactory: TypeAlias = type[sqlite3.Connection] + SqliteRawCursor: TypeAlias = sqlite3.Cursor + +if not TYPE_CHECKING: + SqliteConnection = _SqliteConnection + SqliteConnectionFactory = type[sqlite3.Connection] + SqliteRawCursor = sqlite3.Cursor + __all__ = ( "SqliteConnection", "SqliteConnectionFactory", "SqliteCursor", - "SqliteError", - "SqliteOperationalError", "SqliteRawCursor", "SqliteSessionContext", "sqlite_module", diff --git a/sqlspec/adapters/sqlite/config.py b/sqlspec/adapters/sqlite/config.py index 2916ab93b..8224085b0 100644 --- a/sqlspec/adapters/sqlite/config.py +++ b/sqlspec/adapters/sqlite/config.py @@ -40,7 +40,6 @@ "SqliteConnectionParams", "SqliteDriverFeatures", "SqliteFunctionConfig", - "SqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.sqlite") @@ -89,14 +88,6 @@ class SqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" -class SqliteWindowFunctionConfig(TypedDict): - """User-defined SQLite window function registration.""" - - name: str - narg: int - window_class: "type[Any]" - - class SqliteDriverFeatures(TypedDict): """SQLite driver feature configuration. @@ -126,9 +117,6 @@ class SqliteDriverFeatures(TypedDict): Each entry must include name and func. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. - custom_window_functions: Register user-defined aggregate window functions. - Each entry must include name, narg, and window_class. - default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation. trace_callback: sqlite3 trace hook run for executed statements. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -152,8 +140,6 @@ class SqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[SqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[SqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[SqliteAggregateConfig]]" - custom_window_functions: "NotRequired[Sequence[SqliteWindowFunctionConfig]]" - default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -170,7 +156,6 @@ class SqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", - "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -314,7 +299,6 @@ def get_signature_namespace(self) -> "dict[str, Any]": "SqliteExceptionHandler": SqliteExceptionHandler, "SqliteFunctionConfig": SqliteFunctionConfig, "SqliteSessionContext": SqliteSessionContext, - "SqliteWindowFunctionConfig": SqliteWindowFunctionConfig, }) return namespace @@ -400,9 +384,6 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) - _validate_entries( - runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" - ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/sqlite/core.py b/sqlspec/adapters/sqlite/core.py index 883170a50..47fed0614 100644 --- a/sqlspec/adapters/sqlite/core.py +++ b/sqlspec/adapters/sqlite/core.py @@ -311,14 +311,13 @@ def normalize_execute_parameters(parameters: Any) -> Any: class SqliteStreamSource: """Compiled chunk source streaming dict rows from a SQLite cursor via ``fetchmany``.""" - __slots__ = ("_as_dict", "_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") + __slots__ = ("_chunk_size", "_column_names", "_cursor", "_driver", "_parameters", "_sql") - def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int, as_dict: bool = True) -> None: + def __init__(self, driver: Any, sql: str, parameters: Any, chunk_size: int) -> None: self._driver = driver self._sql = sql self._parameters = parameters self._chunk_size = chunk_size - self._as_dict = as_dict self._cursor: Any = None self._column_names: list[str] | None = None @@ -331,7 +330,7 @@ def start(self) -> None: cursor.execute(self._sql, normalize_execute_parameters(self._parameters)) self._driver._check_pending_exception(handler) - def fetch_chunk(self) -> "list[Any]": + def fetch_chunk(self) -> "list[dict[str, Any]]": handler = self._driver.handle_database_exceptions() rows: list[Any] = [] with handler: @@ -339,8 +338,6 @@ def fetch_chunk(self) -> "list[Any]": self._driver._check_pending_exception(handler) if not rows: return [] - if not self._as_dict: - return rows if self._column_names is None: self._column_names = [description[0] for description in self._cursor.description] return rows_to_dicts(rows, self._column_names) diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index 5e0b17f25..111c443fb 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -1,7 +1,7 @@ """SQLite driver implementation.""" import contextlib -from typing import TYPE_CHECKING, Any, Literal, TypeVar, cast +from typing import TYPE_CHECKING, Any, cast from mypy_extensions import mypyc_attr @@ -54,8 +54,6 @@ __all__ = ("SqliteCursor", "SqliteDriver", "SqliteExceptionHandler", "SqliteSessionContext") -T = TypeVar("T") - @mypyc_attr(allow_interpreted_subclasses=True) class SqliteExceptionHandler(BaseSyncExceptionHandler): @@ -237,21 +235,15 @@ def execute_many( return DMLResult(operation, affected_rows) return super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: + def begin(self) -> None: """Begin a database transaction. - Args: - mode: Transaction lock mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). - Defaults to configured driver feature or SQLite default (DEFERRED). - Raises: SQLSpecError: If transaction cannot be started """ - transaction_mode = mode or self.driver_features.get("default_transaction_mode") try: if not self.connection.in_transaction: - stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" - self.connection.execute(stmt) + self.connection.execute("BEGIN") except sqlite3.Error as e: msg = f"Failed to begin transaction: {e}" raise SQLSpecError(msg) from e @@ -308,14 +300,12 @@ def with_cursor(self, connection: "SqliteConnection") -> "SqliteCursor": """ return SqliteCursor(connection) - def dispatch_select_stream( - self, statement: "SQL", chunk_size: int, as_dict: bool = True - ) -> "SyncRowStream[Any] | None": + def dispatch_select_stream(self, statement: "SQL", chunk_size: int) -> "SyncRowStream[dict[str, Any]] | None": """Return a native SQLite row stream backed by chunked ``fetchmany``.""" if not statement.returns_rows(): return None sql, prepared_parameters = self._compiled_sql(statement, self.statement_config) - return SyncRowStream(SqliteStreamSource(self, sql, prepared_parameters, chunk_size, as_dict=as_dict)) + return SyncRowStream(SqliteStreamSource(self, sql, prepared_parameters, chunk_size)) def handle_database_exceptions(self) -> "SqliteExceptionHandler": """Handle database-specific exceptions and wrap them appropriately. @@ -352,7 +342,6 @@ def load_from_arrow( table: str, source: "ArrowResult | Any", *, - batch_size: int = 10000, partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, @@ -367,12 +356,12 @@ def load_from_arrow( owns_transaction = not self.connection.in_transaction try: if owns_transaction: - self.begin("IMMEDIATE") + self.connection.execute("BEGIN IMMEDIATE") if overwrite: statement = f"DELETE FROM {format_identifier(table)}" with self.with_cursor(self.connection) as cursor: cursor.execute(statement) - for batch in arrow_table.to_batches(max_chunksize=batch_size): + for batch in arrow_table.to_batches(max_chunksize=10000): pydict = batch.to_pydict() records = list(zip(*[pydict[col] for col in columns], strict=False)) if records: @@ -385,10 +374,12 @@ def load_from_arrow( cursor.executemany(insert_sql, cast("Any", prepared_records)) if owns_transaction: self.commit() - except sqlite3.Error as exc: + except BaseException as exc: if owns_transaction: self.rollback() - raise create_mapped_exception(exc) from exc + if isinstance(exc, sqlite3.Error): + raise create_mapped_exception(exc) from exc + raise telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table diff --git a/sqlspec/adapters/sqlite/pool.py b/sqlspec/adapters/sqlite/pool.py index 7a32fd084..aa86fd8d2 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -220,11 +220,7 @@ def _get_thread_connection(self) -> SqliteConnection: last_used = getattr(self._thread_local, "last_used", 0.0) idle_time = now - last_used - if ( - not self._is_memory_db - and idle_time > self._health_check_interval - and not self._is_connection_alive(cast("SqliteConnection", conn)) - ): + if idle_time > self._health_check_interval and not self._is_connection_alive(cast("SqliteConnection", conn)): log_with_context( logger, logging.DEBUG, @@ -362,11 +358,6 @@ def _apply_runtime_setup(connection: SqliteConnection, runtime_setup: "dict[str, aggregate_config["name"], aggregate_config["narg"], aggregate_config["aggregate_class"] ) - create_window_fn = getattr(connection, "create_window_function", None) - if create_window_fn is not None: - for window_config in runtime_setup.get("custom_window_functions", ()): - create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) - for collation_config in runtime_setup.get("custom_collations", ()): connection.create_collation(collation_config["name"], collation_config["func"]) diff --git a/sqlspec/adapters/sqlite/type_converter.py b/sqlspec/adapters/sqlite/type_converter.py index c113759e9..fad74a2e3 100644 --- a/sqlspec/adapters/sqlite/type_converter.py +++ b/sqlspec/adapters/sqlite/type_converter.py @@ -1,3 +1,4 @@ +# Keep in sync with sqlspec/adapters/aiosqlite/type_converter.py """SQLite custom type handlers for optional JSON and type conversion support. Provides registration functions for SQLite's adapter/converter system to enable @@ -8,12 +9,12 @@ instead of lambdas for adapter registration. """ +import json from functools import partial from typing import TYPE_CHECKING, Any from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 from sqlspec.utils.logging import get_logger -from sqlspec.utils.serializers import from_json, to_json if TYPE_CHECKING: from collections.abc import Callable @@ -30,13 +31,14 @@ def json_adapter(value: Any, serializer: "Callable[[Any], str] | None" = None) - Args: value: Python dict or list to serialize. - serializer: Optional JSON serializer callable. Defaults to compiled to_json. + serializer: Optional JSON serializer callable. Defaults to standard json.dumps. Returns: JSON string representation. """ - codec = serializer or to_json - return codec(value) + if serializer is None: + return json.dumps(value, ensure_ascii=False) + return serializer(value) def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = None) -> Any: @@ -44,13 +46,14 @@ def json_converter(value: bytes, deserializer: "Callable[[str], Any] | None" = N Args: value: UTF-8 encoded JSON bytes from SQLite. - deserializer: Optional JSON deserializer callable. Defaults to compiled from_json. + deserializer: Optional JSON deserializer callable. Defaults to standard json.loads. Returns: Deserialized Python object (dict or list). """ - codec = deserializer or from_json - return codec(value.decode("utf-8")) + if deserializer is None: + return json.loads(value.decode("utf-8")) + return deserializer(value.decode("utf-8")) def register_type_handlers( @@ -61,9 +64,6 @@ def register_type_handlers( This function registers handlers globally for the sqlite3 module. It should be called once during application initialization if custom type handling is needed. - Note that sqlite3.register_adapter is deprecated in Python 3.12+ in favor of - statement-level conversions. - Args: json_serializer: Optional custom JSON serializer. json_deserializer: Optional custom JSON deserializer. diff --git a/tests/unit/adapters/test_aiosqlite/test_driver.py b/tests/unit/adapters/test_aiosqlite/test_driver.py index 27fff531c..52f1a340e 100644 --- a/tests/unit/adapters/test_aiosqlite/test_driver.py +++ b/tests/unit/adapters/test_aiosqlite/test_driver.py @@ -401,18 +401,3 @@ async def test_aiosqlite_autocommit_mode_skips_statement_without_open_transactio assert connection.statements == [] assert connection.commit_calls == 0 assert connection.rollback_calls == 0 - - -async def test_aiosqlite_begin_honors_explicit_and_default_transaction_mode() -> None: - """AiosqliteDriver.begin should honor explicit mode and driver_features default_transaction_mode.""" - connection = _AsyncAutocommitConnection(in_transaction=False, autocommit=False) - driver = AiosqliteDriver( - connection=cast("Any", connection), - statement_config=default_statement_config, - driver_features={"default_transaction_mode": "DEFERRED"}, - ) - - await driver.begin() - await driver.begin(mode="EXCLUSIVE") - - assert connection.statements == ["BEGIN DEFERRED", "BEGIN EXCLUSIVE"] diff --git a/tests/unit/adapters/test_aiosqlite/test_pool.py b/tests/unit/adapters/test_aiosqlite/test_pool.py index 6666ca007..cf6d81f70 100644 --- a/tests/unit/adapters/test_aiosqlite/test_pool.py +++ b/tests/unit/adapters/test_aiosqlite/test_pool.py @@ -393,3 +393,17 @@ async def test_wal_setup_failure_propagates_from_new_connection( await pool.new_connection() finally: await pool.close() + + +async def test_close_continues_after_a_native_connection_was_closed() -> None: + pool = AiosqliteConnectionPool({"database": ":memory:"}, pool_size=2) + first = await pool.acquire() + second = await pool.acquire() + try: + await first.connection.close() + await pool.close() + with pytest.raises(ValueError, match="no active connection"): + await second.connection.execute("SELECT 1") + finally: + await second.connection.close() + await pool.close() diff --git a/tests/unit/adapters/test_sqlite/test_driver.py b/tests/unit/adapters/test_sqlite/test_driver.py index 5f0098696..9f045ca0e 100644 --- a/tests/unit/adapters/test_sqlite/test_driver.py +++ b/tests/unit/adapters/test_sqlite/test_driver.py @@ -253,3 +253,33 @@ def test_execute_many_thin_path_checks_all_rows_beyond_sample_threshold() -> Non ) is False ) + + +def test_arrow_ingest_rolls_back_when_a_later_batch_conversion_fails(monkeypatch: pytest.MonkeyPatch) -> None: + import pyarrow as pa + + connection = sqlite3.connect(":memory:") + connection.execute("CREATE TABLE target (value INTEGER)") + driver = SqliteDriver( + connection=connection, driver_features={"storage_capabilities": {"arrow_import_enabled": True}} + ) + batches_prepared = 0 + + def prepare(parameters: Any, *args: Any, **kwargs: Any) -> Any: + nonlocal batches_prepared + batches_prepared += 1 + if batches_prepared == 2: + raise ValueError("conversion failed") + return parameters + + monkeypatch.setattr(SqliteDriver, "_arrow_rows_need_preparation", lambda *args: True) + monkeypatch.setattr( + SqliteDriver, "prepare_driver_parameters", lambda self, *args, **kwargs: prepare(*args, **kwargs) + ) + try: + with pytest.raises(ValueError, match="conversion failed"): + driver.load_from_arrow("target", pa.table({"value": list(range(10001))})) + assert connection.in_transaction is False + assert connection.execute("SELECT COUNT(*) FROM target").fetchone() == (0,) + finally: + connection.close() diff --git a/tests/unit/adapters/test_sqlite/test_pool.py b/tests/unit/adapters/test_sqlite/test_pool.py index 19729d5c7..8b55d89f8 100644 --- a/tests/unit/adapters/test_sqlite/test_pool.py +++ b/tests/unit/adapters/test_sqlite/test_pool.py @@ -378,3 +378,14 @@ def test_enable_wal_reraises_other_operational_errors() -> None: with pytest.raises(sqlite3.OperationalError, match="disk I/O"): pool_module._enable_wal(connection) assert connection.execute.call_count == 1 + + +def test_memory_pool_replaces_a_closed_connection() -> None: + pool = SqliteConnectionPool({"database": ":memory:"}, health_check_interval=-1) + try: + with pool.get_connection() as connection: + connection.close() + with pool.get_connection() as replacement: + assert replacement.execute("SELECT 1").fetchone() == (1,) + finally: + pool.close() diff --git a/tools/scripts/mypyc_inventory.py b/tools/scripts/mypyc_inventory.py index d1fa03953..e618f1cdd 100644 --- a/tools/scripts/mypyc_inventory.py +++ b/tools/scripts/mypyc_inventory.py @@ -234,8 +234,8 @@ "reason": "Async pool runtime compiles without adapter driver exception-handler subclasses.", }, "sqlspec/adapters/sqlite/driver.py": { - "classification": "compile_now", - "reason": "Synchronous SQLite driver compiles cleanly and passes installed-wheel session smoke.", + "classification": "prove_separately", + "reason": "Native compiled driver construction segfaulted in installed-wheel SqliteConfig session smoke; keep interpreted pending driver layout work.", }, "sqlspec/adapters/aiosqlite/driver.py": { "classification": "prove_separately", From 9961dbcb2ef8538a71cc5b794f24a7c30d363e17 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:07:56 +0000 Subject: [PATCH 09/12] test(sqlite): remove checks for reverted feature additions --- tests/integration/adapters/_shared/_driver_type_system.py | 4 ---- tests/unit/adapters/test_aiosqlite/test_config.py | 5 ----- 2 files changed, 9 deletions(-) diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 540f43c86..66bb8a7ca 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -106,8 +106,6 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", - "custom_window_functions", - "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", @@ -259,8 +257,6 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", - "custom_window_functions", - "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", diff --git a/tests/unit/adapters/test_aiosqlite/test_config.py b/tests/unit/adapters/test_aiosqlite/test_config.py index a17c48a2a..4b13f2eb7 100644 --- a/tests/unit/adapters/test_aiosqlite/test_config.py +++ b/tests/unit/adapters/test_aiosqlite/test_config.py @@ -237,11 +237,6 @@ def test_custom_aggregate_entry_requires_keys() -> None: AiosqliteConfig(driver_features={"custom_aggregates": [{"name": "agg", "narg": 1}]}) -def test_custom_window_function_entry_requires_keys() -> None: - with pytest.raises(ImproperConfigurationError, match="custom_window_functions"): - AiosqliteConfig(driver_features={"custom_window_functions": [{"name": "win", "narg": 1}]}) - - def test_progress_handler_interval_must_be_positive() -> None: with pytest.raises(ImproperConfigurationError, match="progress_handler_interval"): AiosqliteConfig(driver_features={"progress_handler": lambda: None, "progress_handler_interval": 0}) From b5af63f6ee52a8a1f7d014b3a77eebccf4c29c2d Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:31:50 +0000 Subject: [PATCH 10/12] fix(sqlite): retain validated native configuration features --- docs/changelog.rst | 7 ++- sqlspec/adapters/aiosqlite/__init__.py | 2 + sqlspec/adapters/aiosqlite/config.py | 19 ++++++++ sqlspec/adapters/aiosqlite/driver.py | 40 +++++++++++++---- sqlspec/adapters/aiosqlite/pool.py | 23 +++++++--- sqlspec/adapters/sqlite/__init__.py | 2 + sqlspec/adapters/sqlite/config.py | 19 ++++++++ sqlspec/adapters/sqlite/driver.py | 21 +++++++-- sqlspec/adapters/sqlite/pool.py | 10 +++++ .../adapters/_shared/_driver_type_system.py | 4 ++ .../adapters/test_aiosqlite/test_driver.py | 12 +++++ .../unit/adapters/test_sqlite/test_driver.py | 44 +++++++++++++++++++ 12 files changed, 181 insertions(+), 22 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index b3a3a64ac..2ad725bed 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,13 +1,10 @@ -========= Changelog -========= All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates -============== Unreleased ---------- @@ -17,6 +14,9 @@ Unreleased * BigQuery supports native query resource controls, explicit STRUCT parameters, typed empty arrays, and configurable Storage Write stream modes while retaining the atomic PENDING default. +* SQLite and aiosqlite support custom window functions on Python 3.11 and + later, explicit transaction locking modes, and configurable Arrow import + batch sizes. Existing transaction defaults remain unchanged. * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. @@ -2277,7 +2277,6 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions -================= For releases before ``v0.24.0``, see the repository tag history and GitHub release records. diff --git a/sqlspec/adapters/aiosqlite/__init__.py b/sqlspec/adapters/aiosqlite/__init__.py index 1f48983b9..daa39414d 100644 --- a/sqlspec/adapters/aiosqlite/__init__.py +++ b/sqlspec/adapters/aiosqlite/__init__.py @@ -7,6 +7,7 @@ AiosqliteDriverFeatures, AiosqliteFunctionConfig, AiosqlitePoolParams, + AiosqliteWindowFunctionConfig, ) from sqlspec.adapters.aiosqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.aiosqlite.driver import AiosqliteDriver, AiosqliteExceptionHandler @@ -34,6 +35,7 @@ "AiosqlitePoolConnection", "AiosqlitePoolParams", "AiosqliteRawCursor", + "AiosqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/aiosqlite/config.py b/sqlspec/adapters/aiosqlite/config.py index 5741c9e7b..d5c80338d 100644 --- a/sqlspec/adapters/aiosqlite/config.py +++ b/sqlspec/adapters/aiosqlite/config.py @@ -48,6 +48,7 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", + "AiosqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -109,6 +110,14 @@ class AiosqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" +class AiosqliteWindowFunctionConfig(TypedDict): + """User-defined aiosqlite window function registration.""" + + name: str + narg: int + window_class: "type[Any]" + + class AiosqliteDriverFeatures(TypedDict): """Aiosqlite driver feature configuration. @@ -138,6 +147,9 @@ class AiosqliteDriverFeatures(TypedDict): Each entry must include name and func. Callable values are plain sync callables. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. + custom_window_functions: Register user-defined aggregate window functions. + Each entry must include name, narg, and window_class. + default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation on the worker thread. trace_callback: sqlite3 trace hook run for executed statements on the worker thread. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -161,6 +173,8 @@ class AiosqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[AiosqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[AiosqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[AiosqliteAggregateConfig]]" + custom_window_functions: "NotRequired[Sequence[AiosqliteWindowFunctionConfig]]" + default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -177,6 +191,7 @@ class AiosqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -348,6 +363,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "AiosqliteFunctionConfig": AiosqliteFunctionConfig, "AiosqlitePoolParams": AiosqlitePoolParams, "AiosqliteSessionContext": AiosqliteSessionContext, + "AiosqliteWindowFunctionConfig": AiosqliteWindowFunctionConfig, "Literal": Literal, "PathLike": PathLike, }) @@ -467,6 +483,9 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) + _validate_entries( + runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" + ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/aiosqlite/driver.py b/sqlspec/adapters/aiosqlite/driver.py index e506fed64..b7ffe5f66 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -4,7 +4,7 @@ import contextlib import inspect import random -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Literal, cast from sqlspec.adapters.aiosqlite._typing import AiosqliteCursor, AiosqliteRawCursor, AiosqliteSessionContext from sqlspec.adapters.aiosqlite._typing import aiosqlite_module as aiosqlite @@ -218,13 +218,26 @@ async def execute_many( await close_result return await super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - async def begin(self) -> None: - """Begin a database transaction.""" + async def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: + """Begin a database transaction. + + Args: + mode: Transaction locking mode (DEFERRED, IMMEDIATE, EXCLUSIVE). + Falls back to ``driver_features['default_transaction_mode']``, + or ``IMMEDIATE`` when neither is set. + """ + transaction_mode = ( + mode if mode is not None else self.driver_features.get("default_transaction_mode", "IMMEDIATE") + ) + stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" + if transaction_mode is not None and transaction_mode not in {"DEFERRED", "IMMEDIATE", "EXCLUSIVE"}: + msg = "Transaction mode must be DEFERRED, IMMEDIATE, or EXCLUSIVE" + raise ValueError(msg) try: if not self.connection.in_transaction: - await self.connection.execute("BEGIN IMMEDIATE") + await self.connection.execute(stmt) except aiosqlite.Error as e: - await _retry_begin_with_backoff(self.connection, e) + await _retry_begin_with_backoff(self.connection, e, statement=stmt) async def commit(self) -> None: """Commit the current transaction.""" @@ -285,11 +298,15 @@ async def load_from_arrow( table: str, source: "ArrowResult | Any", *, + batch_size: int = 10000, partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, ) -> "StorageBridgeJob": """Load Arrow data into SQLite using batched inserts.""" + if isinstance(batch_size, bool) or not isinstance(batch_size, int) or batch_size < 1: + msg = "batch_size must be a positive integer" + raise ValueError(msg) self._require_capability("arrow_import_enabled") arrow_table = self._coerce_arrow_table(source) columns = arrow_table.column_names @@ -304,7 +321,7 @@ async def load_from_arrow( statement = f"DELETE FROM {format_identifier(table)}" async with self.with_cursor(self.connection) as cursor: await cursor.execute(statement) - for batch in arrow_table.to_batches(max_chunksize=10000): + for batch in arrow_table.to_batches(max_chunksize=batch_size): pydict = batch.to_pydict() records = list(zip(*(pydict[col] for col in columns), strict=False)) if records: @@ -546,9 +563,13 @@ def _connection_in_transaction(self) -> bool: async def _retry_begin_with_backoff( - connection: "AiosqliteConnection", initial_error: aiosqlite.Error, max_retries: int = 3 + connection: "AiosqliteConnection", + initial_error: aiosqlite.Error, + max_retries: int = 3, + *, + statement: str = "BEGIN IMMEDIATE", ) -> None: - """Retry ``BEGIN IMMEDIATE`` after SQLite reports a busy connection. + """Retry transaction start 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 @@ -558,6 +579,7 @@ async def _retry_begin_with_backoff( 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. + statement: SQL statement used to start the transaction. Raises: SQLSpecError: If every retry attempt fails. @@ -566,7 +588,7 @@ async def _retry_begin_with_backoff( delay = 0.01 * (2**attempt) + random.uniform(0, 0.01) # noqa: S311 await asyncio.sleep(delay) try: - await connection.execute("BEGIN IMMEDIATE") + await connection.execute(statement) except aiosqlite.Error: if attempt == max_retries - 1: break diff --git a/sqlspec/adapters/aiosqlite/pool.py b/sqlspec/adapters/aiosqlite/pool.py index 67f93bbd2..3325a1650 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -11,7 +11,7 @@ from sqlspec.adapters.aiosqlite._typing import aiosqlite_module as aiosqlite from sqlspec.adapters.aiosqlite._typing import aiosqlite_sqlite_module as sqlite3 from sqlspec.adapters.aiosqlite.core import end_transaction, run_on_worker_thread -from sqlspec.exceptions import SQLSpecError +from sqlspec.exceptions import ImproperConfigurationError, SQLSpecError from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -78,9 +78,12 @@ def _has_active_transaction(connection: "AiosqliteConnection") -> bool: def _register_runtime_objects( - connection: "AiosqliteConnection", aggregates: "Sequence[dict[str, Any]]", collations: "Sequence[dict[str, Any]]" + connection: "AiosqliteConnection", + aggregates: "Sequence[dict[str, Any]]", + collations: "Sequence[dict[str, Any]]", + window_functions: "Sequence[dict[str, Any]]" = (), ) -> None: - """Register custom aggregates and collations on the worker thread.""" + """Register custom aggregates, collations, and window functions on the worker thread.""" raw_connection = connection._conn for aggregate_config in aggregates: raw_connection.create_aggregate( @@ -88,6 +91,13 @@ def _register_runtime_objects( ) for collation_config in collations: raw_connection.create_collation(collation_config["name"], collation_config["func"]) + create_window_fn = getattr(raw_connection, "create_window_function", None) + if window_functions and create_window_fn is None: + msg = "Custom SQLite window functions require Python 3.11 or later" + raise ImproperConfigurationError(msg) + if create_window_fn is not None: + for window_config in window_functions: + create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: "dict[str, Any]") -> None: @@ -115,8 +125,11 @@ async def _apply_runtime_setup(connection: "AiosqliteConnection", runtime_setup: aggregates = runtime_setup.get("custom_aggregates", ()) collations = runtime_setup.get("custom_collations", ()) - if aggregates or collations: - await run_on_worker_thread(connection, _register_runtime_objects, connection, aggregates, collations) + window_functions = runtime_setup.get("custom_window_functions", ()) + if aggregates or collations or window_functions: + await run_on_worker_thread( + connection, _register_runtime_objects, connection, aggregates, collations, window_functions + ) authorizer_callback = runtime_setup.get("authorizer_callback") if authorizer_callback is not None: diff --git a/sqlspec/adapters/sqlite/__init__.py b/sqlspec/adapters/sqlite/__init__.py index 4a17a6019..f534e9894 100644 --- a/sqlspec/adapters/sqlite/__init__.py +++ b/sqlspec/adapters/sqlite/__init__.py @@ -8,6 +8,7 @@ SqliteConnectionParams, SqliteDriverFeatures, SqliteFunctionConfig, + SqliteWindowFunctionConfig, ) from sqlspec.adapters.sqlite.core import build_connection_config, default_statement_config from sqlspec.adapters.sqlite.driver import SqliteDriver, SqliteExceptionHandler @@ -25,6 +26,7 @@ "SqliteDriverFeatures", "SqliteExceptionHandler", "SqliteFunctionConfig", + "SqliteWindowFunctionConfig", "build_connection_config", "default_statement_config", ) diff --git a/sqlspec/adapters/sqlite/config.py b/sqlspec/adapters/sqlite/config.py index 8224085b0..2916ab93b 100644 --- a/sqlspec/adapters/sqlite/config.py +++ b/sqlspec/adapters/sqlite/config.py @@ -40,6 +40,7 @@ "SqliteConnectionParams", "SqliteDriverFeatures", "SqliteFunctionConfig", + "SqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.sqlite") @@ -88,6 +89,14 @@ class SqliteAggregateConfig(TypedDict): aggregate_class: "type[Any]" +class SqliteWindowFunctionConfig(TypedDict): + """User-defined SQLite window function registration.""" + + name: str + narg: int + window_class: "type[Any]" + + class SqliteDriverFeatures(TypedDict): """SQLite driver feature configuration. @@ -117,6 +126,9 @@ class SqliteDriverFeatures(TypedDict): Each entry must include name and func. custom_aggregates: Register SQL aggregates with step/finalize classes. Each entry must include name, narg, and aggregate_class. + custom_window_functions: Register user-defined aggregate window functions. + Each entry must include name, narg, and window_class. + default_transaction_mode: Default SQLite transaction mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). authorizer_callback: sqlite3 authorizer hook run during statement compilation. trace_callback: sqlite3 trace hook run for executed statements. progress_handler: sqlite3 progress hook run every progress_handler_interval VM opcodes. @@ -140,6 +152,8 @@ class SqliteDriverFeatures(TypedDict): custom_functions: "NotRequired[Sequence[SqliteFunctionConfig]]" custom_collations: "NotRequired[Sequence[SqliteCollationConfig]]" custom_aggregates: "NotRequired[Sequence[SqliteAggregateConfig]]" + custom_window_functions: "NotRequired[Sequence[SqliteWindowFunctionConfig]]" + default_transaction_mode: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"]] authorizer_callback: "NotRequired[Callable[[int, str | None, str | None, str | None, str | None], int]]" trace_callback: "NotRequired[Callable[[str], None]]" progress_handler: "NotRequired[Callable[[], int | None]]" @@ -156,6 +170,7 @@ class SqliteDriverFeatures(TypedDict): "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -299,6 +314,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "SqliteExceptionHandler": SqliteExceptionHandler, "SqliteFunctionConfig": SqliteFunctionConfig, "SqliteSessionContext": SqliteSessionContext, + "SqliteWindowFunctionConfig": SqliteWindowFunctionConfig, }) return namespace @@ -384,6 +400,9 @@ def _build_runtime_setup(features: "dict[str, Any]") -> "dict[str, Any] | None": _validate_entries( runtime_setup.get("custom_aggregates", ()), ("name", "narg", "aggregate_class"), "custom_aggregates" ) + _validate_entries( + runtime_setup.get("custom_window_functions", ()), ("name", "narg", "window_class"), "custom_window_functions" + ) interval = runtime_setup.get("progress_handler_interval") if interval is not None and (not isinstance(interval, int) or isinstance(interval, bool) or interval < 1): diff --git a/sqlspec/adapters/sqlite/driver.py b/sqlspec/adapters/sqlite/driver.py index 111c443fb..f159e1d5b 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -1,7 +1,7 @@ """SQLite driver implementation.""" import contextlib -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, Literal, cast from mypy_extensions import mypyc_attr @@ -235,15 +235,24 @@ def execute_many( return DMLResult(operation, affected_rows) return super().execute_many(statement, parameters, *filters, statement_config=statement_config, **kwargs) - def begin(self) -> None: + def begin(self, mode: "Literal['DEFERRED', 'IMMEDIATE', 'EXCLUSIVE'] | None" = None) -> None: """Begin a database transaction. + Args: + mode: Transaction lock mode (DEFERRED, IMMEDIATE, or EXCLUSIVE). + Defaults to configured driver feature or SQLite default (DEFERRED). + Raises: SQLSpecError: If transaction cannot be started """ + transaction_mode = mode if mode is not None else self.driver_features.get("default_transaction_mode") + if transaction_mode is not None and transaction_mode not in {"DEFERRED", "IMMEDIATE", "EXCLUSIVE"}: + msg = "Transaction mode must be DEFERRED, IMMEDIATE, or EXCLUSIVE" + raise ValueError(msg) try: if not self.connection.in_transaction: - self.connection.execute("BEGIN") + stmt = f"BEGIN {transaction_mode}" if transaction_mode else "BEGIN" + self.connection.execute(stmt) except sqlite3.Error as e: msg = f"Failed to begin transaction: {e}" raise SQLSpecError(msg) from e @@ -342,11 +351,15 @@ def load_from_arrow( table: str, source: "ArrowResult | Any", *, + batch_size: int = 10000, partitioner: "dict[str, object] | None" = None, overwrite: bool = False, telemetry: "StorageTelemetry | None" = None, ) -> "StorageBridgeJob": """Load Arrow data into SQLite using chunked batched inserts.""" + if isinstance(batch_size, bool) or not isinstance(batch_size, int) or batch_size < 1: + msg = "batch_size must be a positive integer" + raise ValueError(msg) self._require_capability("arrow_import_enabled") arrow_table = self._coerce_arrow_table(source) columns = arrow_table.column_names @@ -361,7 +374,7 @@ def load_from_arrow( statement = f"DELETE FROM {format_identifier(table)}" with self.with_cursor(self.connection) as cursor: cursor.execute(statement) - for batch in arrow_table.to_batches(max_chunksize=10000): + for batch in arrow_table.to_batches(max_chunksize=batch_size): pydict = batch.to_pydict() records = list(zip(*[pydict[col] for col in columns], strict=False)) if records: diff --git a/sqlspec/adapters/sqlite/pool.py b/sqlspec/adapters/sqlite/pool.py index aa86fd8d2..f88b5fe09 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -10,6 +10,7 @@ from sqlspec.adapters.sqlite._typing import SqliteConnection from sqlspec.adapters.sqlite._typing import sqlite_module as sqlite3 from sqlspec.adapters.sqlite.core import end_transaction +from sqlspec.exceptions import ImproperConfigurationError from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -358,6 +359,15 @@ def _apply_runtime_setup(connection: SqliteConnection, runtime_setup: "dict[str, aggregate_config["name"], aggregate_config["narg"], aggregate_config["aggregate_class"] ) + window_functions = runtime_setup.get("custom_window_functions", ()) + if window_functions: + create_window_fn = getattr(connection, "create_window_function", None) + if create_window_fn is None: + msg = "Custom SQLite window functions require Python 3.11 or later" + raise ImproperConfigurationError(msg) + for window_config in window_functions: + create_window_fn(window_config["name"], window_config["narg"], window_config["window_class"]) + for collation_config in runtime_setup.get("custom_collations", ()): connection.create_collation(collation_config["name"], collation_config["func"]) diff --git a/tests/integration/adapters/_shared/_driver_type_system.py b/tests/integration/adapters/_shared/_driver_type_system.py index 66bb8a7ca..540f43c86 100644 --- a/tests/integration/adapters/_shared/_driver_type_system.py +++ b/tests/integration/adapters/_shared/_driver_type_system.py @@ -106,6 +106,8 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", + "custom_window_functions", + "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", @@ -257,6 +259,8 @@ class SourceEquivalenceCase: "custom_functions", "custom_collations", "custom_aggregates", + "custom_window_functions", + "default_transaction_mode", "authorizer_callback", "trace_callback", "progress_handler", diff --git a/tests/unit/adapters/test_aiosqlite/test_driver.py b/tests/unit/adapters/test_aiosqlite/test_driver.py index 52f1a340e..1c804aabd 100644 --- a/tests/unit/adapters/test_aiosqlite/test_driver.py +++ b/tests/unit/adapters/test_aiosqlite/test_driver.py @@ -401,3 +401,15 @@ async def test_aiosqlite_autocommit_mode_skips_statement_without_open_transactio assert connection.statements == [] assert connection.commit_calls == 0 assert connection.rollback_calls == 0 + + +async def test_aiosqlite_begin_honors_mode_and_rejects_invalid_mode() -> None: + connection = _AsyncAutocommitConnection(in_transaction=False) + driver = AiosqliteDriver( + connection=cast("Any", connection), driver_features={"default_transaction_mode": "EXCLUSIVE"} + ) + await driver.begin(mode="DEFERRED") + await driver.begin() + with pytest.raises(ValueError, match="Transaction mode"): + await driver.begin(mode=cast("Any", "INVALID")) + assert connection.statements == ["BEGIN DEFERRED", "BEGIN EXCLUSIVE"] diff --git a/tests/unit/adapters/test_sqlite/test_driver.py b/tests/unit/adapters/test_sqlite/test_driver.py index 9f045ca0e..7f3f4f114 100644 --- a/tests/unit/adapters/test_sqlite/test_driver.py +++ b/tests/unit/adapters/test_sqlite/test_driver.py @@ -283,3 +283,47 @@ def prepare(parameters: Any, *args: Any, **kwargs: Any) -> Any: assert connection.execute("SELECT COUNT(*) FROM target").fetchone() == (0,) finally: connection.close() + + +def test_sqlite_begin_mode_and_arrow_batch_validation() -> None: + connection = sqlite3.connect(":memory:") + statements: list[str] = [] + connection.set_trace_callback(statements.append) + driver = SqliteDriver(connection=connection, driver_features={"default_transaction_mode": "EXCLUSIVE"}) + try: + driver.begin() + driver.rollback() + driver.begin(mode="DEFERRED") + driver.rollback() + with pytest.raises(ValueError, match="Transaction mode"): + driver.begin(mode=cast("Any", "INVALID")) + with pytest.raises(ValueError, match="batch_size"): + driver.load_from_arrow("sample", cast("Any", None), batch_size=0) + assert statements == ["BEGIN EXCLUSIVE", "ROLLBACK", "BEGIN DEFERRED", "ROLLBACK"] + finally: + connection.close() + + +@pytest.mark.parametrize("async_adapter", [False, True]) +def test_window_registration_checks_native_capability(async_adapter: bool) -> None: + from types import SimpleNamespace + from unittest.mock import Mock + + from sqlspec.adapters.aiosqlite.pool import _register_runtime_objects + from sqlspec.adapters.sqlite.pool import _apply_runtime_setup + from sqlspec.exceptions import ImproperConfigurationError + + window_class = type("Window", (), {}) + functions = [{"name": "window_sum", "narg": 1, "window_class": window_class}] + register = Mock() + native = SimpleNamespace(create_window_function=register) + if async_adapter: + _register_runtime_objects(cast("Any", SimpleNamespace(_conn=native)), (), (), functions) + else: + _apply_runtime_setup(cast("Any", native), {"custom_window_functions": functions}) + register.assert_called_once_with("window_sum", 1, window_class) + with pytest.raises(ImproperConfigurationError, match="Python 3"): + if async_adapter: + _register_runtime_objects(cast("Any", SimpleNamespace(_conn=object())), (), (), functions) + else: + _apply_runtime_setup(cast("Any", object()), {"custom_window_functions": functions}) From ab99efc8d568ee8871a6fdc7aff20d4ded2201c9 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 21:34:06 +0000 Subject: [PATCH 11/12] test(sqlite): retain window configuration validation coverage --- tests/unit/adapters/test_aiosqlite/test_config.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/unit/adapters/test_aiosqlite/test_config.py b/tests/unit/adapters/test_aiosqlite/test_config.py index 4b13f2eb7..a17c48a2a 100644 --- a/tests/unit/adapters/test_aiosqlite/test_config.py +++ b/tests/unit/adapters/test_aiosqlite/test_config.py @@ -237,6 +237,11 @@ def test_custom_aggregate_entry_requires_keys() -> None: AiosqliteConfig(driver_features={"custom_aggregates": [{"name": "agg", "narg": 1}]}) +def test_custom_window_function_entry_requires_keys() -> None: + with pytest.raises(ImproperConfigurationError, match="custom_window_functions"): + AiosqliteConfig(driver_features={"custom_window_functions": [{"name": "win", "narg": 1}]}) + + def test_progress_handler_interval_must_be_positive() -> None: with pytest.raises(ImproperConfigurationError, match="progress_handler_interval"): AiosqliteConfig(driver_features={"progress_handler": lambda: None, "progress_handler_interval": 0}) From ae5a930dd422c7be9efef5df200489968d037496 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Sun, 27 Sep 2026 22:05:47 +0000 Subject: [PATCH 12/12] docs: reconcile unreleased adapter changelog after rebase --- docs/changelog.rst | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/docs/changelog.rst b/docs/changelog.rst index 2ad725bed..8769e74ec 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -1,10 +1,13 @@ +========= Changelog +========= All notable SQLSpec changes are summarized here. Entries are grouped by release and focus on user-visible behavior, public API changes, compatibility notes, and important operational fixes. Recent Updates +============== Unreleased ---------- @@ -14,9 +17,9 @@ Unreleased * BigQuery supports native query resource controls, explicit STRUCT parameters, typed empty arrays, and configurable Storage Write stream modes while retaining the atomic PENDING default. -* SQLite and aiosqlite support custom window functions on Python 3.11 and - later, explicit transaction locking modes, and configurable Arrow import - batch sizes. Existing transaction defaults remain unchanged. +* SQLite and aiosqlite can register custom window functions on Python 3.11 + and later when the SQLite runtime supports them. Choose a transaction lock + mode or set the batch size for Arrow imports. Defaults stay the same. * Arrow ODBC runs ``execute_many()`` one row at a time. It reports an unknown row count since the native driver does not return the number of changed rows. @@ -45,6 +48,9 @@ Unreleased **Fixed:** +* SQLite pools replace lost in-memory connections. Arrow imports roll back + writes on failure or cancellation when the adapter owns the transaction. + * Builder results keep CTE trees independent, and column pruning no longer exposes its cached expression to mutation. SQL generation avoids redundant copies of temporary trees while preserving caller and cache ownership. @@ -2277,6 +2283,7 @@ v0.24.0 - Builder consolidation * Refactored builder code to reduce duplication. Previous Versions +================= For releases before ``v0.24.0``, see the repository tag history and GitHub release records.