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

"""Psqlpy 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.psqlpy.config import PsqlpyConfig


__all__ = ("PsqlpyLitestarConfig", "PsqlpyStore")


[docs] class PsqlpyLitestarConfig(LitestarConfig): """Psqlpy-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 PsqlpyStore(BaseSQLSpecStore["PsqlpyConfig"]): """PostgreSQL session store using Psqlpy driver. Implements server-side session storage for Litestar using PostgreSQL via the Psqlpy driver (Rust-based async driver). Provides efficient session management with: - Native async PostgreSQL operations via Rust - UPSERT support using ON CONFLICT - Automatic expiration handling - Efficient cleanup of expired sessions Args: config: PsqlpyConfig instance. """ __slots__ = () extension_config_options = BaseSQLSpecStore.extension_config_options | frozenset({ "autovacuum_analyze_scale_factor", "autovacuum_vacuum_scale_factor", "fillfactor", }) def __init__(self, config: "PsqlpyConfig") -> None: """Initialize Psqlpy session store. Args: config: PsqlpyConfig 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. """ new_expires_at = self._calculate_expires_at(renew_for) if renew_for is not None else None if new_expires_at is not None: 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: query_result = await conn.fetch(sql, [new_expires_at, key]) rows = query_result.result() if query_result else [] if not rows: return None return bytes(rows[0]["data"]) sql = f""" SELECT data 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: query_result = await conn.fetch(sql, [key]) rows = query_result.result() if query_result else [] if not rows: return None return bytes(rows[0]["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: query_result = await conn.fetch(sql, [key]) rows = query_result.result() return len(rows) > 0 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: query_result = await conn.fetch(sql, [key]) rows = query_result.result() if not rows: return None expires_at = rows[0]["expires_at"] 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 RETURNING session_id """ async with self._config.provide_connection() as conn: query_result = await conn.fetch(sql, []) rows = query_result.result() count = len(rows) 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]