"""Base session store classes for Litestar integration."""
import re
from abc import abstractmethod
from datetime import datetime, timedelta, timezone
from typing import Any, ClassVar, Final, Generic, TypeVar, cast
from litestar.stores.base import Store
from sqlspec.exceptions import ImproperConfigurationError
from sqlspec.migrations.schema import SchemaTarget, ensure_schema_async, ensure_schema_sync
from sqlspec.observability import resolve_db_system
from sqlspec.utils.logging import get_logger
from sqlspec.utils.sync_tools import async_
from sqlspec.utils.type_guards import has_extension_config
__all__ = ("BaseSQLSpecStore",)
ConfigT = TypeVar("ConfigT")
logger = get_logger("sqlspec.extensions.litestar.store")
VALID_TABLE_NAME_PATTERN: Final = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
MAX_TABLE_NAME_LENGTH: Final = 63
[docs]
class BaseSQLSpecStore(Store, Generic[ConfigT]):
"""Base class for SQLSpec-backed Litestar session stores.
Inherits Litestar's ``Store`` abstract base class for server-side session
storage backed by SQLSpec database adapters.
This abstract base class provides common functionality for all database-specific
store implementations including:
- Connection management via SQLSpec configs
- Session expiration calculation
- Table creation utilities
Subclasses must implement dialect-specific SQL queries.
Args:
config: SQLSpec database configuration with extension_config["litestar"] settings.
"""
__slots__ = ("_config", "_table_name")
extension_config_options: ClassVar[frozenset[str]] = frozenset({
"auto_trace_headers",
"commit_mode",
"connection_key",
"correlation_header",
"correlation_headers",
"create_schema",
"disable_di",
"enable_correlation_middleware",
"enable_sqlcommenter_middleware",
"extra_commit_statuses",
"extra_rollback_statuses",
"manage_schema",
"migrations_path",
"pool_key",
"run_migrations",
"session_key",
"session_table",
})
[docs]
def __init__(self, config: ConfigT) -> None:
"""Initialize the session store.
Args:
config: SQLSpec database configuration.
"""
self._config = config
self._validate_extension_config()
self._table_name = self._table_name_from_config()
self._ensure_table_name(self._table_name)
@property
def config(self) -> ConfigT:
"""Return the database configuration."""
return self._config
@property
def table_name(self) -> str:
"""Return the session table name."""
return self._table_name
@property
def create_schema_enabled(self) -> bool:
"""Return whether adapter-level creation should run."""
manage_schema, create_schema = self._schema_management_flags()
return manage_schema and create_schema
[docs]
@abstractmethod
async def delete_expired(self) -> int:
"""Delete all expired sessions.
Returns:
Number of sessions deleted.
"""
raise NotImplementedError
[docs]
@abstractmethod
async def create_table(self) -> None:
"""Create the session table if it doesn't exist."""
raise NotImplementedError
[docs]
def prepare_schema_sync(self, driver: Any) -> None:
"""Prepare adapter-specific schema decisions with a synchronous driver."""
[docs]
async def prepare_schema_async(self, driver: Any) -> None:
"""Prepare adapter-specific schema decisions with an asynchronous driver."""
[docs]
async def reconcile_schema(self, *, assume_existing: bool = False) -> None:
"""Apply additive session-table changes from the canonical store DDL.
Args:
assume_existing: Skip table discovery after adapter-level creation.
"""
manage_schema, create_schema = self._schema_management_flags()
if not manage_schema:
return
statement_config = getattr(self._config, "statement_config", None)
dialect = getattr(statement_config, "dialect", None)
target = SchemaTarget.from_ddl(self._table_name, self._table_ddl(), dialect=dialect)
session_context = cast("Any", self._config).provide_session()
if hasattr(session_context, "__aenter__"):
async with session_context as driver:
await ensure_schema_async(
driver, [target], manage_schema=True, create_schema=create_schema, assume_existing=assume_existing
)
return
await async_(_reconcile_schema_sync)(
self._config, target, create_schema=create_schema, assume_existing=assume_existing
)
def _table_name_from_config(self) -> str:
"""Extract the configured session table name."""
default_name = "litestar_session"
if has_extension_config(self._config):
extension_config = cast("dict[str, dict[str, Any]]", self._config.extension_config)
litestar_config: dict[str, Any] = extension_config.get("litestar", {})
session_table = litestar_config.get("session_table", default_name)
if session_table is True:
return default_name
return str(session_table)
return default_name
def _validate_extension_config(self) -> None:
"""Reject Litestar options that this adapter store cannot honor."""
if not has_extension_config(self._config):
return
extension_config = cast("dict[str, dict[str, Any]]", self._config.extension_config)
settings = extension_config.get("litestar", {})
unsupported = sorted(set(settings).difference(type(self).extension_config_options))
if unsupported:
adapter = type(self).__module__.split(".")[2]
keys = ", ".join(repr(key) for key in unsupported)
msg = f"Unsupported Litestar configuration key(s) for {adapter}: {keys}"
raise ImproperConfigurationError(msg)
def _schema_management_flags(self) -> "tuple[bool, bool]":
"""Return automatic-management and missing-table creation flags."""
if not has_extension_config(self._config):
return True, True
extension_config = cast("dict[str, dict[str, Any]]", self._config.extension_config)
settings = extension_config.get("litestar", {})
return bool(settings.get("manage_schema", True)), bool(settings.get("create_schema", True))
@abstractmethod
def _table_ddl(self) -> str:
"""Get the CREATE TABLE SQL for this database dialect.
Returns:
SQL statement to create the sessions table.
"""
raise NotImplementedError
@abstractmethod
def _drop_table_sql(self) -> "list[str]":
"""Get the DROP TABLE SQL statements for this database dialect.
Returns:
List of SQL statements to drop the table and all indexes.
Order matters: drop indexes before table.
"""
raise NotImplementedError
def _log_table_created(self) -> None:
logger.debug(
"Litestar session table ready",
extra={"db.system": resolve_db_system(type(self).__name__), "session_table": self._table_name},
)
def _log_delete_all(self) -> None:
logger.debug(
"Litestar sessions cleared",
extra={"db.system": resolve_db_system(type(self).__name__), "session_table": self._table_name},
)
def _log_delete_expired(self, count: int) -> None:
logger.debug(
"Litestar sessions expired cleanup",
extra={
"db.system": resolve_db_system(type(self).__name__),
"session_table": self._table_name,
"deleted_sessions": count,
},
)
def _calculate_expires_at(self, expires_in: "int | timedelta | None") -> "datetime | None":
"""Calculate expiration timestamp from expires_in.
Args:
expires_in: Seconds or timedelta until expiration.
Returns:
UTC datetime of expiration, or None if no expiration.
"""
if expires_in is None:
return None
expires_in_seconds = int(expires_in.total_seconds()) if isinstance(expires_in, timedelta) else expires_in
if expires_in_seconds <= 0:
return None
return datetime.now(timezone.utc) + timedelta(seconds=expires_in_seconds)
def _value_to_bytes(self, value: "str | bytes") -> bytes:
"""Convert value to bytes if needed.
Args:
value: String or bytes value.
Returns:
Value as bytes.
"""
if isinstance(value, str):
return value.encode("utf-8")
return value
@staticmethod
def _ensure_table_name(table_name: str) -> None:
"""Validate table name for SQL safety.
Args:
table_name: Table name to validate.
Raises:
ValueError: If table name is invalid.
"""
if not table_name:
msg = "Table name cannot be empty"
raise ValueError(msg)
if len(table_name) > MAX_TABLE_NAME_LENGTH:
msg = f"Table name too long: {len(table_name)} chars (max {MAX_TABLE_NAME_LENGTH})"
raise ValueError(msg)
if not VALID_TABLE_NAME_PATTERN.match(table_name):
msg = f"Invalid table name: {table_name!r}. Must start with letter/underscore and contain only alphanumeric characters and underscores"
raise ValueError(msg)
def _reconcile_schema_sync(config: Any, target: SchemaTarget, *, create_schema: bool, assume_existing: bool) -> None:
"""Run additive reconciliation for a synchronous Litestar store."""
with config.provide_session() as driver:
ensure_schema_sync(
driver, [target], manage_schema=True, create_schema=create_schema, assume_existing=assume_existing
)