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