diff --git a/docs/changelog.rst b/docs/changelog.rst index b3a3a64ac..8769e74ec 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -17,6 +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 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. 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 f1b94435d..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.config 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: @@ -1036,7 +1039,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..d5c80338d 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,6 +48,7 @@ "AiosqliteDriverFeatures", "AiosqliteFunctionConfig", "AiosqlitePoolParams", + "AiosqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.aiosqlite") @@ -106,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. @@ -135,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. @@ -158,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]]" @@ -168,14 +185,13 @@ 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", "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -184,12 +200,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): @@ -353,6 +363,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "AiosqliteFunctionConfig": AiosqliteFunctionConfig, "AiosqlitePoolParams": AiosqlitePoolParams, "AiosqliteSessionContext": AiosqliteSessionContext, + "AiosqliteWindowFunctionConfig": AiosqliteWindowFunctionConfig, "Literal": Literal, "PathLike": PathLike, }) @@ -440,54 +451,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 - - 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 +468,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): @@ -520,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 5141e1321..ecb8a7e75 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 @@ -41,8 +43,10 @@ _T = TypeVar("_T") __all__ = ( + "SQLITE_CONNECT_SUPPORTS_AUTOCOMMIT", "AiosqliteStreamSource", "apply_driver_features", + "apply_extension_pragmas", "build_connection_config", "build_insert_statement", "build_profile", @@ -51,12 +55,18 @@ "create_mapped_exception", "default_statement_config", "driver_profile", + "end_transaction", + "execute_and_resolve_metadata", "execute_and_resolve_rowcount", "execute_fetchall_with_description", + "execute_fetchall_with_metadata", + "execute_many_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", @@ -66,6 +76,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 @@ -81,20 +99,82 @@ 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: """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) @@ -368,8 +448,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 +455,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 +515,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, @@ -482,6 +558,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, @@ -491,7 +615,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 +638,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 +652,25 @@ 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() + + +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 + + def _resolve_insert_target(expression: Any) -> "tuple[str | None, str] | None": if not isinstance(expression, exp.Insert): return None @@ -689,10 +832,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..b7ffe5f66 100644 --- a/sqlspec/adapters/aiosqlite/driver.py +++ b/sqlspec/adapters/aiosqlite/driver.py @@ -1,23 +1,26 @@ """AIOSQLite driver implementation for async SQLite operations.""" import asyncio +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 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, + end_transaction, + execute_and_resolve_metadata, + execute_fetchall_with_metadata, + execute_many_on_worker_thread, format_identifier, - normalize_execute_many_parameters, normalize_execute_parameters, resolve_rowcount, run_on_worker_thread, @@ -54,6 +57,9 @@ "AiosqliteSessionContext", ) +_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. @@ -147,12 +153,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": @@ -194,30 +200,49 @@ 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: + 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: - """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.""" 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 @@ -225,7 +250,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 @@ -273,38 +298,48 @@ 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, 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() - except (aiosqlite.Error, sqlite3.Error) as exc: + await self.commit() + except BaseException as exc: if owns_transaction: - await self.connection.rollback() - raise create_mapped_exception(exc) from exc + await self.rollback() + 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 @@ -528,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 @@ -540,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. @@ -548,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/events/store.py b/sqlspec/adapters/aiosqlite/events/store.py index e843c78f5..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 @@ -36,11 +37,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..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.config 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 @@ -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: @@ -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: @@ -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,12 +218,15 @@ 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 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 35695418d..3325a1650 100644 --- a/sqlspec/adapters/aiosqlite/pool.py +++ b/sqlspec/adapters/aiosqlite/pool.py @@ -10,13 +10,13 @@ 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.exceptions import SQLSpecError +from sqlspec.adapters.aiosqlite.core import end_transaction, run_on_worker_thread +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 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 @@ -40,16 +40,22 @@ 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 +77,29 @@ 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]]", + window_functions: "Sequence[dict[str, Any]]" = (), +) -> None: + """Register custom aggregates, collations, and window functions 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"]) + 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: pragmas = runtime_setup.get("pragmas", ()) if pragmas: @@ -94,19 +123,12 @@ 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", ()): + 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, - 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"] + connection, _register_runtime_objects, connection, aggregates, collations, window_functions ) authorizer_callback = runtime_setup.get("authorizer_callback") @@ -219,7 +241,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.""" @@ -228,10 +250,9 @@ 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: - # 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 +292,7 @@ class AiosqliteConnectionPool: """Multi-connection pool for aiosqlite.""" __slots__ = ( - "_closed_event_instance", + "_closed_event", "_connect_timeout", "_connection_parameters", "_connection_registry", @@ -279,13 +300,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 +354,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 +367,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 +376,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 +397,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.""" + with suppress(Exception): + raw_conn = getattr(connection.connection, "_conn", None) + if raw_conn is not None: + 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 +476,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: @@ -645,8 +654,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 +751,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 +800,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,11 +816,9 @@ 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() + await end_transaction(connection.connection, commit=False) connection.idle_since = time.time() self._queue.put_nowait(connection) except Exception as e: @@ -854,6 +852,11 @@ async def close(self) -> None: self._connection_registry.clear() if connections: + for conn in connections: + 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/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/adk/store.py b/sqlspec/adapters/sqlite/adk/store.py index 9297a883d..4c22f5ffd 100644 --- a/sqlspec/adapters/sqlite/adk/store.py +++ b/sqlspec/adapters/sqlite/adk/store.py @@ -8,7 +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, render_pragmas from sqlspec.config import ADKConfig from sqlspec.exceptions import ImproperConfigurationError from sqlspec.extensions.adk import BaseSyncADKStore, StoredEvent, StoredSession, normalize_session_list_options @@ -84,7 +84,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 @@ -112,7 +111,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) @@ -135,7 +133,7 @@ def create_session( 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 +153,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: @@ -181,7 +178,7 @@ def get_session( 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 +207,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) @@ -223,7 +219,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)) - conn.commit() + end_transaction(conn, commit=True) def list_sessions( self, @@ -286,13 +282,12 @@ 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 +295,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"]) @@ -324,7 +318,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 +346,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)) @@ -410,13 +403,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." @@ -451,7 +444,6 @@ def get_events( Returns: List of event records ordered by timestamp ASC. """ - """Synchronous implementation of get_events.""" if limit == 0: return [] @@ -497,7 +489,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: @@ -509,7 +500,7 @@ def delete_expired_events(self, before: datetime, app_name: "str | None" = None) 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): @@ -529,7 +520,7 @@ def delete_idle_sessions(self, updated_before: datetime, app_name: "str | None" 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 +529,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: @@ -550,7 +540,7 @@ def delete_idle_user_states(self, updated_before: datetime, app_name: "str | Non 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,7 +549,6 @@ 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: @@ -575,7 +564,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} @@ -595,7 +583,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 (?, ?, ?) @@ -607,11 +594,10 @@ 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 (?, ?, ?, ?) @@ -623,11 +609,10 @@ 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: @@ -643,7 +628,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 (?, ?) @@ -653,7 +637,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)) - conn.commit() + end_transaction(conn, commit=True) def _apply_pragmas(self, connection: Any) -> None: """Apply PRAGMA optimization profile for this connection. @@ -798,7 +782,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. @@ -815,7 +798,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 @@ -838,26 +820,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 +852,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 +879,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 +906,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 +918,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 +939,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 +1101,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..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 @@ -37,6 +40,7 @@ "SqliteConnectionParams", "SqliteDriverFeatures", "SqliteFunctionConfig", + "SqliteWindowFunctionConfig", ) logger = get_logger("sqlspec.adapters.sqlite") @@ -85,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. @@ -114,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. @@ -137,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]]" @@ -147,14 +164,13 @@ 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", "custom_aggregates", "custom_collations", "custom_functions", + "custom_window_functions", "extensions", "pragmas", "progress_handler", @@ -163,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): @@ -251,7 +261,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 +314,7 @@ def get_signature_namespace(self) -> "dict[str, Any]": "SqliteExceptionHandler": SqliteExceptionHandler, "SqliteFunctionConfig": SqliteFunctionConfig, "SqliteSessionContext": SqliteSessionContext, + "SqliteWindowFunctionConfig": SqliteWindowFunctionConfig, }) return namespace @@ -358,54 +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 - - 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 +385,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 +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/core.py b/sqlspec/adapters/sqlite/core.py index 39a7e56e9..47fed0614 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 @@ -36,11 +37,14 @@ 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", + "apply_extension_pragmas", "build_connection_config", "build_insert_statement", "build_profile", @@ -49,10 +53,13 @@ "create_mapped_exception", "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", @@ -77,6 +84,89 @@ 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( + 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() + + +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 @@ -294,13 +384,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 +408,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 +448,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 +455,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 +480,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 +489,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 +519,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, @@ -644,10 +723,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..f159e1d5b 100644 --- a/sqlspec/adapters/sqlite/driver.py +++ b/sqlspec/adapters/sqlite/driver.py @@ -1,6 +1,9 @@ """SQLite driver implementation.""" -from typing import TYPE_CHECKING, Any, cast +import contextlib +from typing import TYPE_CHECKING, Any, Literal, 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 @@ -19,32 +22,40 @@ 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") +@mypyc_attr(allow_interpreted_subclasses=True) class SqliteExceptionHandler(BaseSyncExceptionHandler): """Context manager for handling SQLite database exceptions. @@ -67,6 +78,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. @@ -75,7 +87,7 @@ class SqliteDriver(SyncDriverAdapterBase): """ __slots__ = ("_data_dictionary", "_rowid_target_cache") - dialect = "sqlite" + dialect: "DialectType | None" = "sqlite" def __init__( self, @@ -207,27 +219,40 @@ 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) 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 @@ -239,7 +264,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. @@ -312,7 +337,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,20 +351,21 @@ 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.""" + 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, 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: @@ -349,16 +374,25 @@ def load_from_arrow( 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() - except sqlite3.Error as exc: + self.commit() + except BaseException as exc: if owns_transaction: - self.connection.rollback() - raise create_mapped_exception(exc) from exc + self.rollback() + if isinstance(exc, sqlite3.Error): + raise create_mapped_exception(exc) from exc + raise telemetry_payload = self._ingest_telemetry(arrow_table) telemetry_payload["destination"] = table @@ -406,6 +440,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) @@ -430,14 +465,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) @@ -455,10 +492,13 @@ 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" - 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"}: @@ -493,7 +533,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]) @@ -506,7 +546,6 @@ 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]] if row_len == 1: if has_type_coercion and coercion_map is not None: for param_set in parameters: diff --git a/sqlspec/adapters/sqlite/events/store.py b/sqlspec/adapters/sqlite/events/store.py index 1b5233d10..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 @@ -36,11 +37,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..b63892b0b 100644 --- a/sqlspec/adapters/sqlite/litestar/store.py +++ b/sqlspec/adapters/sqlite/litestar/store.py @@ -5,7 +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 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_ @@ -63,7 +63,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 +75,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 +201,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 +210,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 +232,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 +249,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 +257,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 +265,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 +273,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 +311,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..f88b5fe09 100644 --- a/sqlspec/adapters/sqlite/pool.py +++ b/sqlspec/adapters/sqlite/pool.py @@ -9,7 +9,8 @@ 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.exceptions import ImproperConfigurationError from sqlspec.utils.logging import POOL_LOGGER_NAME, get_logger, log_with_context from sqlspec.utils.uuids import uuid4 @@ -28,47 +29,25 @@ 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() - - class SqliteConnectionPool: """Thread-local connection manager for SQLite. @@ -84,6 +63,7 @@ class SqliteConnectionPool: "_enable_optimizations", "_generation", "_health_check_interval", + "_is_memory_db", "_on_connection_create", "_pool_id", "_recycle_seconds", @@ -116,6 +96,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 @@ -156,10 +138,7 @@ def new_connection(self) -> SqliteConnection: try: if self._enable_optimizations: - database = self._connection_parameters.get("database", ":memory:") - is_memory = database == ":memory:" or "mode=memory" in str(database) - - if is_memory: + if self._is_memory_db: connection.execute("PRAGMA journal_mode = MEMORY") connection.execute("PRAGMA synchronous = OFF") connection.execute("PRAGMA temp_store = MEMORY") @@ -183,7 +162,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 +181,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 +212,16 @@ 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 idle_time > self._health_check_interval and not self._is_connection_alive(cast("SqliteConnection", conn)): log_with_context( logger, logging.DEBUG, @@ -245,12 +232,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 +251,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]": @@ -282,11 +270,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.""" @@ -316,12 +304,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).""" @@ -374,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_config.py b/tests/unit/adapters/test_aiosqlite/test_config.py index 7cd6c43d6..a17c48a2a 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)) @@ -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..1c804aabd 100644 --- a/tests/unit/adapters/test_aiosqlite/test_driver.py +++ b/tests/unit/adapters/test_aiosqlite/test_driver.py @@ -345,3 +345,71 @@ 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_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_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 d8b37891e..7f3f4f114 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,150 @@ 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 + ) + + +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() + + +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}) diff --git a/tests/unit/adapters/test_sqlite/test_pool.py b/tests/unit/adapters/test_sqlite/test_pool.py index 289c34f61..8b55d89f8 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 @@ -377,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()