import enum from collections.abc import Callable from datetime import datetime from typing import Any, ClassVar import sqlalchemy as sa from fansifter_common.adapters.db import Database from fansifter_common.adapters.db.models import ModelMixin from fansifter_common.utils.functional import lazy_proxy from sqlalchemy import event from sqlalchemy.orm import DeclarativeBase, MappedAsDataclass from sqlalchemy.pool import QueuePool from app.adapters import aws_dsql from app.config import settings def get_db() -> Database: url = settings.db_url if settings.dsql_endpoint: # Inject a fresh token into the URL so the pool can connect on startup, # before do_connect is registered. url = sa.engine.make_url(url).set( password=aws_dsql.generate_db_connect_admin_auth_token() ) database = Database( url=url, engine_args={ "echo": settings.db_echo, "poolclass": QueuePool, "pool_size": settings.db_pool_size, "max_overflow": settings.db_pool_max_overflow, "pool_recycle": settings.db_pool_recycle, "pool_pre_ping": True, "skip_autocommit_rollback": True, "connect_args": { "keepalives": 1, "keepalives_idle": 60, "keepalives_interval": 10, "keepalives_count": 5, }, "use_native_hstore": False, }, ) if settings.dsql_endpoint: event.listen(database.engine, "do_connect", _dsql_token_refresher()) return database def _dsql_token_refresher() -> Callable[..., None]: """Returns a do_connect listener that injects a fresh DSQL token as password. Called on every new physical connection, so the token is always valid regardless of pool_recycle or connection drops. """ def refresh( _dialect: Any, _conn_rec: Any, _cargs: list[Any], cparams: dict[str, Any], ) -> None: cparams["password"] = aws_dsql.generate_db_connect_admin_auth_token() return refresh db = lazy_proxy(get_db) class Model(ModelMixin, MappedAsDataclass, DeclarativeBase): __db__ = db type_annotation_map: ClassVar[dict[Any, Any]] = { enum.Enum: sa.Enum(length=32, native_enum=False), datetime: sa.DateTime(timezone=True), }