Source code for sqlspec.adapters.sqlite.config

"""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_extensions import NotRequired

from sqlspec.adapters.sqlite._typing import (
    SqliteConnection,
    SqliteConnectionFactory,
    SqliteCursor,
    SqliteSessionContext,
)
from sqlspec.adapters.sqlite.core import apply_driver_features, build_connection_config, default_statement_config
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
from sqlspec.config import ExtensionConfigs, SyncDatabaseConfig
from sqlspec.driver._sync import SyncPoolConnectionContext, SyncPoolSessionFactory
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.utils.logging import get_logger
from sqlspec.utils.uuids import uuid4

if TYPE_CHECKING:
    from collections.abc import Callable, Sequence

    from sqlspec.core import StatementConfig
    from sqlspec.observability import ObservabilityConfig

__all__ = (
    "SqliteAggregateConfig",
    "SqliteCollationConfig",
    "SqliteConfig",
    "SqliteConnectionParams",
    "SqliteDriverFeatures",
    "SqliteFunctionConfig",
)

logger = get_logger("sqlspec.adapters.sqlite")


class SqliteConnectionParams(TypedDict):
    """SQLite connection parameters."""

    database: NotRequired[str | PathLike[str]]
    timeout: NotRequired[float]
    detect_types: NotRequired[int]
    isolation_level: NotRequired[Literal["DEFERRED", "IMMEDIATE", "EXCLUSIVE"] | None]
    check_same_thread: NotRequired[bool]
    factory: "NotRequired[SqliteConnectionFactory | None]"
    cached_statements: NotRequired[int]
    uri: NotRequired[bool]
    autocommit: NotRequired[bool]
    pool_recycle_seconds: NotRequired[int]
    health_check_interval: NotRequired[float]
    enable_optimizations: NotRequired[bool]
    enable_foreign_keys: NotRequired[bool]
    extra: NotRequired[dict[str, Any]]


class SqliteFunctionConfig(TypedDict):
    """User-defined SQLite function registration."""

    name: str
    narg: int
    func: "Callable[..., Any]"
    deterministic: NotRequired[bool]


class SqliteCollationConfig(TypedDict):
    """User-defined SQLite collation registration."""

    name: str
    func: "Callable[[str, str], int]"


class SqliteAggregateConfig(TypedDict):
    """User-defined SQLite aggregate registration."""

    name: str
    narg: int
    aggregate_class: "type[Any]"


class SqliteDriverFeatures(TypedDict):
    """SQLite driver feature configuration.

    Controls optional type handling and serialization features for SQLite connections.

    enable_custom_adapters: Enable custom type adapters for JSON/UUID/datetime conversion.
     Defaults to True for enhanced Python type support.
     Set to False only if you need pure SQLite behavior without type conversions.
    json_serializer: Custom JSON serializer function.
     Defaults to sqlspec.utils.serializers.to_json.
    json_deserializer: Custom JSON deserializer function.
     Defaults to sqlspec.utils.serializers.from_json.
    on_connection_create: Callback executed when a connection is created.
     Receives the raw sqlite3 connection for low-level driver configuration.
     Runs after internal setup (PRAGMA optimizations).
    enable_events: Enable database event channel support.
     Defaults to True when extension_config["events"] is configured.
     Provides pub/sub capabilities via table-backed queue (SQLite has no native pub/sub).
     Requires extension_config["events"] for migration setup.
    events_backend: Event channel backend selection.
     Only option: "poll_queue" (durable table-backed queue with lease-based retries and acknowledgements).
     SQLite does not have native pub/sub, so poll_queue is the only backend.
     Defaults to "poll_queue".
    custom_functions: Register SQL functions that run on the connection thread.
     Each entry must include name, narg, and func.
    custom_collations: Register SQL collations that compare two string values.
     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.
    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.
    progress_handler_interval: Progress callback interval in SQLite virtual machine opcodes.
     Must be a positive integer when provided.
    row_factory: Row factory selector or callable used for raw sqlite3 connections.
     "row" maps to sqlite3.Row, "dict" maps to a dict row adapter, "tuple" keeps tuple rows.
     "dict" and custom callables can change raw connection result shapes seen by callers.
    text_factory: Text factory used for raw sqlite3 connections.
    pragmas: Additional PRAGMA settings applied after built-in optimization PRAGMAs.
     User values override built-in defaults when the same PRAGMA appears in both places.
    extensions: Shared-library extension paths loaded on each connection.
    """

    enable_custom_adapters: NotRequired[bool]
    json_serializer: "NotRequired[Callable[[Any], str]]"
    json_deserializer: "NotRequired[Callable[[str], Any]]"
    on_connection_create: "NotRequired[Callable[[SqliteConnection], None]]"
    enable_events: NotRequired[bool]
    events_backend: NotRequired[Literal["poll_queue"]]
    custom_functions: "NotRequired[Sequence[SqliteFunctionConfig]]"
    custom_collations: "NotRequired[Sequence[SqliteCollationConfig]]"
    custom_aggregates: "NotRequired[Sequence[SqliteAggregateConfig]]"
    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]]"
    progress_handler_interval: NotRequired[int]
    row_factory: "NotRequired[Literal['row', 'dict', 'tuple'] | Callable[..., Any]]"
    text_factory: "NotRequired[Callable[[bytes], Any]]"
    pragmas: "NotRequired[Mapping[str, str | int | bool]]"
    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",
    "extensions",
    "pragmas",
    "progress_handler",
    "progress_handler_interval",
    "row_factory",
    "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):
    """Context manager for Sqlite connections."""

    __slots__ = ()


class _SqliteSessionConnectionHandler(SyncPoolSessionFactory):
    __slots__ = ()


class SqliteConfig(SyncDatabaseConfig[SqliteConnection, SqliteConnectionPool, SqliteDriver]):
    """SQLite configuration with thread-local connections."""

    driver_type: "ClassVar[type[SqliteDriver]]" = SqliteDriver
    connection_type: "ClassVar[type[SqliteConnection]]" = SqliteConnection
    supports_transactional_ddl: "ClassVar[bool]" = True
    supports_native_arrow_export: "ClassVar[bool]" = True
    supports_native_arrow_import: "ClassVar[bool]" = True
    supports_native_parquet_export: "ClassVar[bool]" = True
    supports_native_parquet_import: "ClassVar[bool]" = True
    supports_native_row_streaming: "ClassVar[bool]" = True
    _connection_context_class: "ClassVar[type[SqliteConnectionContext]]" = SqliteConnectionContext
    _session_factory_class: "ClassVar[type[_SqliteSessionConnectionHandler]]" = _SqliteSessionConnectionHandler
    _session_context_class: "ClassVar[type[SqliteSessionContext]]" = SqliteSessionContext
    _default_statement_config = default_statement_config

    def __init__(
        self,
        *,
        connection_config: "SqliteConnectionParams | dict[str, Any] | None" = None,
        connection_instance: "SqliteConnectionPool | None" = None,
        migration_config: "dict[str, Any] | None" = None,
        statement_config: "StatementConfig | None" = None,
        driver_features: "SqliteDriverFeatures | dict[str, Any] | None" = None,
        bind_key: "str | None" = None,
        extension_config: "ExtensionConfigs | None" = None,
        observability_config: "ObservabilityConfig | None" = None,
        **kwargs: Any,
    ) -> None:
        """Initialize SQLite configuration.

        Args:
            connection_config: Configuration parameters including connection settings
            connection_instance: Pre-created pool instance
            migration_config: Migration configuration
            statement_config: Default SQL statement configuration
            driver_features: Optional driver feature configuration
            bind_key: Optional bind key for the configuration
            extension_config: Extension-specific configuration
            observability_config: Adapter-level observability overrides for lifecycle hooks and observers
            **kwargs: Additional keyword arguments passed to the base configuration.
        """
        config_dict: dict[str, Any] = dict(connection_config) if connection_config else {}
        if "database" not in config_dict or config_dict["database"] == ":memory:":
            config_dict["database"] = f"file:memory_{uuid4().hex}?mode=memory&cache=private"
            config_dict["uri"] = True
        elif "database" in config_dict:
            database_path = str(config_dict["database"])
            if database_path.startswith("file:") and not config_dict.get("uri"):
                logger.debug(
                    "Database URI detected (%s) but uri=True not set. "
                    "Auto-enabling URI mode to prevent physical file creation.",
                    database_path,
                )
                config_dict["uri"] = True

        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
        )
        self._runtime_setup: dict[str, Any] | None = _build_runtime_setup(features_dict)

        super().__init__(
            bind_key=bind_key,
            connection_instance=connection_instance,
            connection_config=config_dict,
            migration_config=migration_config,
            statement_config=statement_config,
            driver_features=features_dict,
            extension_config=extension_config,
            observability_config=observability_config,
            **kwargs,
        )
def create_connection(self) -> SqliteConnection: """Get a SQLite connection from the pool. Returns: SqliteConnection: A connection from the pool """ pool = self.provide_pool() return pool.acquire()