"""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]