Source code for sqlspec.adapters.asyncpg.litestar.store

"""AsyncPG session store for Litestar integration."""

from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, cast

from typing_extensions import NotRequired

from sqlspec.config import LitestarConfig
from sqlspec.extensions.litestar.store import BaseSQLSpecStore

if TYPE_CHECKING:
    from sqlspec.adapters.asyncpg.config import AsyncpgConfig


__all__ = ("AsyncpgLitestarConfig", "AsyncpgStore")


[docs] class AsyncpgLitestarConfig(LitestarConfig): """Asyncpg-specific Litestar settings. Use inside ``extension_config["litestar"]`` with this adapter's session store. """ fillfactor: NotRequired[int] """Table fillfactor. Default: 80.""" autovacuum_vacuum_scale_factor: NotRequired[float] """Table autovacuum vacuum scale factor.""" autovacuum_analyze_scale_factor: NotRequired[float] """Table autovacuum analyze scale factor."""
class AsyncpgStore(BaseSQLSpecStore["AsyncpgConfig"]): """PostgreSQL session store using AsyncPG driver. Implements server-side session storage for Litestar using PostgreSQL via the AsyncPG driver. Provides efficient session management with: - Native async PostgreSQL operations - UPSERT support using ON CONFLICT - Automatic expiration handling - Efficient cleanup of expired sessions Args: config: AsyncpgConfig instance with extension_config["litestar"] settings. """ __slots__ = () extension_config_options = BaseSQLSpecStore.extension_config_options | frozenset({ "autovacuum_analyze_scale_factor", "autovacuum_vacuum_scale_factor", "fillfactor", }) def __init__(self, config: "AsyncpgConfig") -> None: """Initialize AsyncPG session store. Args: config: AsyncpgConfig instance. """ super().__init__(config) async def create_table(self) -> None: """Create the session table if it doesn't exist.""" if not self.create_schema_enabled: await self.reconcile_schema() return sql = self._table_ddl() async with self._config.provide_session() as driver: await driver.execute_script(sql) self._log_table_created() await self.reconcile_schema(assume_existing=True) async def get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Get a session value by key. Args: key: Session ID to retrieve. renew_for: If given, renew the expiry time for this duration. Returns: Session data as bytes if found and not expired, None otherwise. """ if renew_for is not None: new_expires_at = self._calculate_expires_at(renew_for) if new_expires_at is not None: update_sql = f""" UPDATE {self._table_name} SET expires_at = CASE WHEN expires_at IS NOT NULL THEN $1 ELSE expires_at END, updated_at = CURRENT_TIMESTAMP WHERE session_id = $2 AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) RETURNING data """ async with self._config.provide_connection() as conn: row = await conn.fetchrow(update_sql, new_expires_at, key) if row is None: return None return bytes(row["data"]) sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = $1 AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ async with self._config.provide_connection() as conn: row = await conn.fetchrow(sql, key) if row is None: return None return bytes(row["data"]) async def set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Store a session value. Args: key: Session ID. value: Session data. expires_in: Time until expiration. """ data = self._value_to_bytes(value) expires_at = self._calculate_expires_at(expires_in) sql = f""" INSERT INTO {self._table_name} (session_id, data, expires_at) VALUES ($1, $2, $3) ON CONFLICT (session_id) DO UPDATE SET data = EXCLUDED.data, expires_at = EXCLUDED.expires_at, updated_at = CURRENT_TIMESTAMP """ async with self._config.provide_connection() as conn: await conn.execute(sql, key, data, expires_at) async def delete(self, key: str) -> None: """Delete a session by key. Args: key: Session ID to delete. """ sql = f"DELETE FROM {self._table_name} WHERE session_id = $1" async with self._config.provide_connection() as conn: await conn.execute(sql, key) async def delete_all(self) -> None: """Delete all sessions from the store.""" sql = f"DELETE FROM {self._table_name}" async with self._config.provide_connection() as conn: await conn.execute(sql) self._log_delete_all() async def exists(self, key: str) -> bool: """Check if a session key exists and is not expired. Args: key: Session ID to check. Returns: True if the session exists and is not expired. """ sql = f""" SELECT 1 FROM {self._table_name} WHERE session_id = $1 AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ async with self._config.provide_connection() as conn: result = await conn.fetchval(sql, key) return result is not None async def expires_in(self, key: str) -> "int | None": """Get the time in seconds until the session expires. Args: key: Session ID to check. Returns: Seconds until expiration, or None if no expiry or key doesn't exist. """ sql = f""" SELECT expires_at FROM {self._table_name} WHERE session_id = $1 """ async with self._config.provide_connection() as conn: expires_at = await conn.fetchval(sql, key) if expires_at is None: return None now = datetime.now(timezone.utc) if expires_at <= now: return 0 delta = expires_at - now return int(delta.total_seconds()) async def delete_expired(self) -> int: """Delete all expired sessions. Returns: Number of sessions deleted. """ sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP" async with self._config.provide_connection() as conn: result = await conn.execute(sql) count = int(result.split()[-1]) if count > 0: self._log_delete_expired(count) return count def _table_ddl(self) -> str: """Get PostgreSQL CREATE TABLE SQL with optimized schema. Returns: SQL statement to create the sessions table with proper indexes. """ fillfactor, vacuum_scale, analyze_scale = _postgres_litestar_tuning(self._config) return f""" CREATE TABLE IF NOT EXISTS {self._table_name} ( session_id TEXT PRIMARY KEY, data BYTEA NOT NULL, expires_at TIMESTAMPTZ, created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP ) WITH (fillfactor = {fillfactor}); CREATE INDEX IF NOT EXISTS idx_{self._table_name}_expires_at ON {self._table_name}(expires_at) WHERE expires_at IS NOT NULL; ALTER TABLE {self._table_name} SET ( autovacuum_vacuum_scale_factor = {vacuum_scale:g}, autovacuum_analyze_scale_factor = {analyze_scale:g} ); """ def _drop_table_sql(self) -> "list[str]": """Get PostgreSQL DROP TABLE SQL statements. Returns: List of SQL statements to drop indexes and table. """ return [f"DROP INDEX IF EXISTS idx_{self._table_name}_expires_at", f"DROP TABLE IF EXISTS {self._table_name}"] def _postgres_litestar_tuning(config: Any) -> "tuple[int, float, float]": settings = cast("dict[str, Any]", config.extension_config.get("litestar", {})) fillfactor = settings.get("fillfactor", 80) if not isinstance(fillfactor, int) or isinstance(fillfactor, bool) or fillfactor not in range(10, 101): msg = "extension_config['litestar']['fillfactor'] must be an integer from 10 to 100" raise ValueError(msg) scales: list[float] = [] for key, default in (("autovacuum_vacuum_scale_factor", 0.05), ("autovacuum_analyze_scale_factor", 0.02)): value = settings.get(key, default) if not isinstance(value, (int, float)) or isinstance(value, bool) or not 0 <= float(value) <= 1: msg = f"extension_config['litestar']['{key}'] must be a number from 0 to 1" raise ValueError(msg) scales.append(float(value)) return fillfactor, scales[0], scales[1]