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

"""Psycopg session stores for Litestar integration.

Provides both async and sync PostgreSQL session stores using psycopg3.
"""

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

from typing_extensions import NotRequired

from sqlspec.adapters.psycopg._typing import psycopg_dict_row as dict_row
from sqlspec.config import LitestarConfig
from sqlspec.extensions.litestar.store import BaseSQLSpecStore
from sqlspec.utils.sync_tools import async_

if TYPE_CHECKING:
    from sqlspec.adapters.psycopg.config import PsycopgAsyncConfig, PsycopgSyncConfig


__all__ = ("PsycopgAsyncStore", "PsycopgLitestarConfig", "PsycopgSyncStore")


[docs] class PsycopgLitestarConfig(LitestarConfig): """Psycopg-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 PsycopgAsyncStore(BaseSQLSpecStore["PsycopgAsyncConfig"]): """PostgreSQL session store using Psycopg async driver. Implements server-side session storage for Litestar using PostgreSQL via the Psycopg (psycopg3) async 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: PsycopgAsyncConfig instance. """ __slots__ = () extension_config_options = BaseSQLSpecStore.extension_config_options | frozenset({ "autovacuum_analyze_scale_factor", "autovacuum_vacuum_scale_factor", "fillfactor", }) def __init__(self, config: "PsycopgAsyncConfig") -> None: """Initialize Psycopg async session store. Args: config: PsycopgAsyncConfig 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) await driver.commit() 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. """ sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = %s AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ conn_context = self._config.provide_connection() async with conn_context as conn: async with conn.cursor(row_factory=dict_row) as cur: await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() if row is None: return None if renew_for is not None and row["expires_at"] 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 = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ await conn.execute(update_sql.encode(), (new_expires_at, key)) await conn.commit() 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 (%s, %s, %s) ON CONFLICT (session_id) DO UPDATE SET data = EXCLUDED.data, expires_at = EXCLUDED.expires_at, updated_at = CURRENT_TIMESTAMP """ conn_context = self._config.provide_connection() async with conn_context as conn: await conn.execute(sql.encode(), (key, data, expires_at)) await conn.commit() 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 = %s" conn_context = self._config.provide_connection() async with conn_context as conn: await conn.execute(sql.encode(), (key,)) await conn.commit() async def delete_all(self) -> None: """Delete all sessions from the store.""" sql = f"DELETE FROM {self._table_name}" conn_context = self._config.provide_connection() async with conn_context as conn: await conn.execute(sql.encode()) await conn.commit() 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 = %s AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: await cur.execute(sql.encode(), (key,)) result = await cur.fetchone() 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 = %s """ conn_context = self._config.provide_connection() async with conn_context as conn: async with conn.cursor(row_factory=dict_row) as cur: await cur.execute(sql.encode(), (key,)) row = await cur.fetchone() if row is None or row["expires_at"] is None: return None expires_at = row["expires_at"] 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" conn_context = self._config.provide_connection() async with conn_context as conn, conn.cursor() as cur: await cur.execute(sql.encode()) await conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 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}"] class PsycopgSyncStore(BaseSQLSpecStore["PsycopgSyncConfig"]): """PostgreSQL session store using Psycopg sync driver. Implements server-side session storage for Litestar using PostgreSQL via the synchronous Psycopg (psycopg3) driver. Uses Litestar's sync_to_thread utility to provide an async interface compatible with the Store protocol. Provides efficient session management with: - Sync operations wrapped for async compatibility - UPSERT support using ON CONFLICT - Automatic expiration handling - Efficient cleanup of expired sessions Args: config: PsycopgSyncConfig instance. """ __slots__ = () extension_config_options = BaseSQLSpecStore.extension_config_options | frozenset({ "autovacuum_analyze_scale_factor", "autovacuum_vacuum_scale_factor", "fillfactor", }) def __init__(self, config: "PsycopgSyncConfig") -> None: """Initialize Psycopg sync session store. Args: config: PsycopgSyncConfig 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 await async_(self._create_table)() 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. """ return await async_(self._get)(key, renew_for) 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. """ await async_(self._set)(key, value, expires_in) async def delete(self, key: str) -> None: """Delete a session by key. Args: key: Session ID to delete. """ await async_(self._delete)(key) async def delete_all(self) -> None: """Delete all sessions from the store.""" await async_(self._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. """ return await async_(self._exists)(key) 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. """ return await async_(self._expires_in)(key) async def delete_expired(self) -> int: """Delete all expired sessions. Returns: Number of sessions deleted. """ return await async_(self._delete_expired)() 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 _create_table(self) -> None: """Synchronous implementation of create_table.""" sql = self._table_ddl() with self._config.provide_session() as driver: driver.execute_script(sql) driver.commit() self._log_table_created() def _get(self, key: str, renew_for: "int | timedelta | None" = None) -> "bytes | None": """Synchronous implementation of get.""" sql = f""" SELECT data, expires_at FROM {self._table_name} WHERE session_id = %s AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: cur.execute(sql.encode(), (key,)) row = cur.fetchone() if row is None: return None if renew_for is not None and row["expires_at"] 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 = %s, updated_at = CURRENT_TIMESTAMP WHERE session_id = %s """ conn.execute(update_sql.encode(), (new_expires_at, key)) conn.commit() return bytes(row["data"]) def _set(self, key: str, value: "str | bytes", expires_in: "int | timedelta | None" = None) -> None: """Synchronous implementation of set.""" 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 (%s, %s, %s) ON CONFLICT (session_id) DO UPDATE SET data = EXCLUDED.data, expires_at = EXCLUDED.expires_at, updated_at = CURRENT_TIMESTAMP """ with self._config.provide_connection() as conn: conn.execute(sql.encode(), (key, data, expires_at)) conn.commit() def _delete(self, key: str) -> None: """Synchronous implementation of delete.""" sql = f"DELETE FROM {self._table_name} WHERE session_id = %s" with self._config.provide_connection() as conn: conn.execute(sql.encode(), (key,)) conn.commit() def _delete_all(self) -> None: """Synchronous implementation of delete_all.""" sql = f"DELETE FROM {self._table_name}" with self._config.provide_connection() as conn: conn.execute(sql.encode()) conn.commit() self._log_delete_all() def _exists(self, key: str) -> bool: """Synchronous implementation of exists.""" sql = f""" SELECT 1 FROM {self._table_name} WHERE session_id = %s AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP) """ with self._config.provide_connection() as conn, conn.cursor() as cur: cur.execute(sql.encode(), (key,)) result = cur.fetchone() return result is not None def _expires_in(self, key: str) -> "int | None": """Synchronous implementation of expires_in.""" sql = f""" SELECT expires_at FROM {self._table_name} WHERE session_id = %s """ with self._config.provide_connection() as conn: with conn.cursor(row_factory=dict_row) as cur: cur.execute(sql.encode(), (key,)) row = cur.fetchone() if row is None or row["expires_at"] is None: return None expires_at = row["expires_at"] now = datetime.now(timezone.utc) if expires_at <= now: return 0 delta = expires_at - now return int(delta.total_seconds()) def _delete_expired(self) -> int: """Synchronous implementation of delete_expired.""" sql = f"DELETE FROM {self._table_name} WHERE expires_at <= CURRENT_TIMESTAMP" with self._config.provide_connection() as conn, conn.cursor() as cur: cur.execute(sql.encode()) conn.commit() count = cur.rowcount if cur.rowcount and cur.rowcount > 0 else 0 if count > 0: self._log_delete_expired(count) return count 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]