"""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)