Source code for sqlspec.extensions.starlette.extension

import logging
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Any

from sqlspec.base import SQLSpec
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.extensions.starlette._state import SQLSpecConfigState
from sqlspec.extensions.starlette._utils import get_or_create_session, get_state_value
from sqlspec.extensions.starlette.middleware import (
    CorrelationMiddleware,
    SQLCommenterMiddleware,
    SQLSpecAutocommitMiddleware,
    SQLSpecManualMiddleware,
)
from sqlspec.utils.logging import get_logger, log_with_context
from sqlspec.utils.sync_tools import ensure_async_

if TYPE_CHECKING:
    from collections.abc import AsyncGenerator

    from starlette.applications import Starlette
    from starlette.requests import Request

__all__ = ("SQLSpecPlugin",)

logger = get_logger("sqlspec.extensions.starlette")

DEFAULT_COMMIT_MODE = "manual"
DEFAULT_CONNECTION_KEY = "db_connection"
DEFAULT_POOL_KEY = "db_pool"
DEFAULT_SESSION_KEY = "db_session"


class SQLSpecPlugin:
    """SQLSpec integration for Starlette applications.

    Provides middleware-based session management, automatic transaction handling,
    and connection pooling lifecycle management.
    """

    __slots__ = ("_config_states", "_correlation_middleware_added", "_sqlcommenter_middleware_added", "_sqlspec")

    def __init__(self, sqlspec: SQLSpec, app: "Starlette | None" = None) -> None:
        """Initialize SQLSpec Starlette extension.

        Args:
            sqlspec: Pre-configured SQLSpec instance with registered configs.
            app: Optional Starlette application to initialize immediately.
        """
        self._sqlspec = sqlspec
        self._config_states: list[SQLSpecConfigState] = []
        self._correlation_middleware_added = False
        self._sqlcommenter_middleware_added = False

        for cfg in self._sqlspec.configs.values():
            settings = self._extract_extension_settings(cfg)
            state = self._config_state(cfg, settings)
            self._config_states.append(state)

        if app is not None:
            self.init_app(app)
        log_with_context(
            logger,
            logging.DEBUG,
            "extension.init",
            framework="starlette",
            stage="init",
            config_count=len(self._config_states),
        )
def _extract_extension_settings(self, config: Any) -> "dict[str, Any]": """Extract Starlette settings from config.extension_config. Args: config: Database configuration instance. Returns: Dictionary of Starlette-specific settings. """ framework_config = config.extension_config.get("starlette", {}) pool_key = framework_config.get("pool_key", DEFAULT_POOL_KEY) if not config.supports_connection_pooling and pool_key == DEFAULT_POOL_KEY: pool_key = f"_{DEFAULT_POOL_KEY}_{id(config)}" correlation_headers = framework_config.get("correlation_headers") return { "connection_key": framework_config.get("connection_key", DEFAULT_CONNECTION_KEY), "pool_key": pool_key, "session_key": framework_config.get("session_key", DEFAULT_SESSION_KEY), "commit_mode": framework_config.get("commit_mode", DEFAULT_COMMIT_MODE), "extra_commit_statuses": framework_config.get("extra_commit_statuses"), "extra_rollback_statuses": framework_config.get("extra_rollback_statuses"), "disable_di": framework_config.get("disable_di", False), "enable_correlation_middleware": framework_config.get("enable_correlation_middleware", False), "correlation_header": framework_config.get("correlation_header", "x-request-id"), "correlation_headers": tuple(correlation_headers) if correlation_headers is not None else None, "auto_trace_headers": framework_config.get("auto_trace_headers", True), "enable_sqlcommenter_middleware": framework_config.get("enable_sqlcommenter_middleware", True), "sqlcommenter_framework": framework_config.get("sqlcommenter_framework", "starlette"), } def _config_state(self, config: Any, settings: "dict[str, Any]") -> SQLSpecConfigState: """Create configuration state object. Args: config: Database configuration instance. settings: Extracted framework settings. Returns: Configuration state instance. """ return SQLSpecConfigState( config=config, connection_key=settings["connection_key"], pool_key=settings["pool_key"], session_key=settings["session_key"], commit_mode=settings["commit_mode"], extra_commit_statuses=settings["extra_commit_statuses"], extra_rollback_statuses=settings["extra_rollback_statuses"], disable_di=settings["disable_di"], enable_correlation_middleware=settings["enable_correlation_middleware"], correlation_header=settings["correlation_header"], correlation_headers=settings["correlation_headers"], auto_trace_headers=settings["auto_trace_headers"], enable_sqlcommenter_middleware=settings["enable_sqlcommenter_middleware"], sqlcommenter_framework=settings["sqlcommenter_framework"], ) def init_app(self, app: "Starlette") -> None: """Initialize Starlette application with SQLSpec. Validates configuration, wraps lifespan, and adds middleware. Args: app: Starlette application instance. """ self._ensure_unique_keys() original_lifespan = app.router.lifespan_context @asynccontextmanager async def combined_lifespan(app: "Starlette") -> "AsyncGenerator[None, None]": async with self.lifespan(app), original_lifespan(app): yield app.router.lifespan_context = combined_lifespan for config_state in self._config_states: if not config_state.disable_di: self._add_middleware(app, config_state) # Add correlation middleware if any config enables it (only add once) self._add_correlation_middleware(app) self._add_sqlcommenter_middleware(app) log_with_context( logger, logging.DEBUG, "extension.init", framework="starlette", stage="configured", config_count=len(self._config_states), )