Source code for sqlspec.extensions.sanic.extension

import logging
from typing import TYPE_CHECKING, Any

from sqlspec.base import SQLSpec
from sqlspec.core import CorrelationExtractor
from sqlspec.core.sqlcommenter import SQLCommenterContext
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.extensions.sanic._state import SanicConfigState
from sqlspec.extensions.sanic._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.sync_tools import ensure_async_, with_ensure_async_
from sqlspec.utils.type_guards import has_name

if TYPE_CHECKING:
    from sanic import Sanic

__all__ = ("SQLSpecPlugin",)

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

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


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

    Provides Sanic-native configuration parsing and request helper methods.
    Runtime listener and middleware behavior is registered by ``init_app``.
    """

    __slots__ = ("_config_states", "_extractor", "_lifecycle_listeners_added", "_request_middleware_added", "_sqlspec")

    def __init__(self, sqlspec: SQLSpec, app: "Sanic[Any, Any] | None" = None) -> None:
        """Initialize SQLSpec Sanic extension.

        Args:
            sqlspec: Pre-configured SQLSpec instance with registered configs.
            app: Optional Sanic application to initialize immediately.
        """
        self._sqlspec = sqlspec
        self._config_states: list[SanicConfigState] = []
        self._lifecycle_listeners_added = False
        self._request_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)

        correlation_state = self._first_correlation_state()
        self._extractor = (
            CorrelationExtractor(
                primary_header=correlation_state.correlation_header,
                additional_headers=correlation_state.correlation_headers,
                auto_trace_headers=correlation_state.auto_trace_headers,
            )
            if correlation_state is not None
            else None
        )

        if app is not None:
            self.init_app(app)

        log_with_context(
            logger,
            logging.DEBUG,
            "extension.init",
            framework="sanic",
            stage="init",
            config_count=len(self._config_states),
        )
def _extract_extension_settings(self, config: Any) -> "dict[str, Any]": """Extract Sanic settings from config.extension_config. Args: config: Database configuration instance. Returns: Dictionary of Sanic-specific settings. """ framework_config = config.extension_config.get("sanic", {}) 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", "sanic"), } def _config_state(self, config: Any, settings: "dict[str, Any]") -> SanicConfigState: """Create configuration state object. Args: config: Database configuration instance. settings: Extracted Sanic settings. Returns: Configuration state instance. """ return SanicConfigState( 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: "Sanic[Any, Any]") -> None: """Initialize Sanic application with SQLSpec. Args: app: Sanic application instance. """ self._ensure_unique_keys() setattr(app.ctx, "sqlspec_plugin", self) self._add_lifecycle_listeners(app) self._add_request_middleware(app)