Source code for sqlspec.extensions.flask.extension

"""Flask extension for SQLSpec database integration."""

import atexit
import logging
from typing import TYPE_CHECKING, Any, Literal

from sqlspec.base import SQLSpec
from sqlspec.config import AsyncDatabaseConfig, NoPoolAsyncConfig
from sqlspec.core import CorrelationExtractor
from sqlspec.core.sqlcommenter import SQLCommenterContext
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.extensions.flask._state import FlaskConfigState
from sqlspec.extensions.flask._utils import (
    get_context_value,
    get_or_create_session,
    has_context_value,
    pop_context_value,
    set_context_value,
)
from sqlspec.utils.correlation import CorrelationContext
from sqlspec.utils.logging import get_logger, log_with_context
from sqlspec.utils.portal import PortalProvider

if TYPE_CHECKING:
    from flask import Flask, Response

__all__ = ("SQLSpecPlugin",)

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

DEFAULT_COMMIT_MODE: Literal["manual"] = "manual"
DEFAULT_SESSION_KEY = "db_session"


class SQLSpecPlugin:
    """Flask extension for SQLSpec database integration.

    Provides request-scoped session management, automatic transaction handling,
    and async adapter support via portal pattern.
    """

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

        Args:
            sqlspec: SQLSpec instance with registered configs.
            app: Optional Flask application to initialize immediately.
        """
        self._sqlspec = sqlspec
        self._config_states: list[FlaskConfigState] = []
        self._portal: PortalProvider | None = None
        self._has_async_configs = False
        self._cleanup_registered = False
        self._shutdown_complete = False
        self._enable_correlation = False
        self._enable_sqlcommenter = False
        self._extractor: CorrelationExtractor | None = None

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

            if state.is_async:
                self._has_async_configs = True

            if state.enable_correlation_middleware and not self._enable_correlation:
                self._enable_correlation = True
                self._extractor = CorrelationExtractor(
                    primary_header=state.correlation_header,
                    additional_headers=state.correlation_headers,
                    auto_trace_headers=state.auto_trace_headers,
                )
            if (
                state.enable_sqlcommenter_middleware
                and state.config.statement_config.enable_sqlcommenter
                and not self._enable_sqlcommenter
            ):
                self._enable_sqlcommenter = True

        if app is not None:
            self.init_app(app)
def _config_state(self, config: Any) -> FlaskConfigState: """Create configuration state from database config. Args: config: Database configuration instance. Returns: FlaskConfigState instance. """ flask_config = config.extension_config.get("flask", {}) session_key = flask_config.get("session_key", DEFAULT_SESSION_KEY) connection_key = flask_config.get("connection_key", f"sqlspec_connection_{session_key}") commit_mode = flask_config.get("commit_mode", DEFAULT_COMMIT_MODE) extra_commit_statuses = flask_config.get("extra_commit_statuses") extra_rollback_statuses = flask_config.get("extra_rollback_statuses") disable_di = flask_config.get("disable_di", False) enable_correlation = flask_config.get("enable_correlation_middleware", False) correlation_header = flask_config.get("correlation_header", "x-request-id") correlation_headers = flask_config.get("correlation_headers") if correlation_headers is not None: correlation_headers = tuple(correlation_headers) auto_trace_headers = flask_config.get("auto_trace_headers", True) enable_sqlcommenter = flask_config.get("enable_sqlcommenter_middleware", True) is_async = isinstance(config, (AsyncDatabaseConfig, NoPoolAsyncConfig)) return FlaskConfigState( config=config, connection_key=connection_key, session_key=session_key, commit_mode=commit_mode, extra_commit_statuses=extra_commit_statuses, extra_rollback_statuses=extra_rollback_statuses, is_async=is_async, disable_di=disable_di, enable_correlation_middleware=enable_correlation, correlation_header=correlation_header, correlation_headers=correlation_headers, auto_trace_headers=auto_trace_headers, enable_sqlcommenter_middleware=enable_sqlcommenter, ) def init_app(self, app: "Flask") -> None: """Initialize Flask application with SQLSpec. Validates configuration, creates portal if needed, creates pools, and registers hooks. Args: app: Flask application to initialize. Raises: ImproperConfigurationError: If extension already registered or keys not unique. """ if "sqlspec" in app.extensions: msg = "SQLSpec extension already registered on this Flask application" raise ImproperConfigurationError(msg) self._ensure_unique_keys() if self._has_async_configs: self._portal = PortalProvider() self._portal.start() log_with_context(logger, logging.DEBUG, "extension.init", framework="flask", stage="portal_started") pools: dict[str, Any] = {} for config_state in self._config_states: if config_state.config.supports_connection_pooling: if config_state.is_async: pool = self._portal.portal.call(config_state.config.create_pool) # type: ignore[union-attr,arg-type] else: pool = config_state.config.create_pool() pools[config_state.session_key] = pool log_with_context( logger, logging.DEBUG, "session.create", framework="flask", session_key=config_state.session_key ) app.extensions["sqlspec"] = {"plugin": self, "pools": pools} if any(not state.disable_di for state in self._config_states): app.before_request(self._before_request_handler) app.after_request(self._after_request_handler) app.teardown_appcontext(self._teardown_appcontext_handler) self._register_shutdown_hook() log_with_context( logger, logging.DEBUG, "extension.init", framework="flask", stage="configured", config_count=len(self._config_states), async_enabled=self._has_async_configs, )