import enum import functools import logging import random import time from collections.abc import Callable from datetime import datetime from typing import Any, ClassVar, overload import sqlalchemy as sa from fansifter_common.adapters.db import Database as _Database from fansifter_common.adapters.db.base import _Transaction from fansifter_common.adapters.db.models import ModelMixin from fansifter_common.utils.functional import lazy_proxy from sqlalchemy import event from sqlalchemy.exc import OperationalError from sqlalchemy.orm import DeclarativeBase, MappedAsDataclass from sqlalchemy.pool import QueuePool from resonance_engine.adapters import aws_dsql from resonance_engine.config import settings logger = logging.getLogger(__name__) _CONFLICT_TYPES = ("SerializationFailure", "TransactionConflict") _MAX_RETRIES = 5 _BACKOFF_S = 0.1 def is_conflict(exc: BaseException) -> bool: orig = getattr(exc, "orig", None) return orig is not None and type(orig).__name__ in _CONFLICT_TYPES def _retry[**P, R](func: Callable[P, R]) -> Callable[P, R]: """Wrap *func* with exponential-backoff retry on DSQL serialization conflicts.""" @functools.wraps(func) def wrapper(*args: P.args, **kwargs: P.kwargs) -> R: for attempt in range(_MAX_RETRIES + 1): try: return func(*args, **kwargs) except OperationalError as exc: if is_conflict(exc) and attempt < _MAX_RETRIES: delay = _BACKOFF_S * (2**attempt) * (0.5 + random.random()) logger.warning( "DB conflict on attempt %d/%d — retrying in %.2fs", attempt + 1, _MAX_RETRIES, delay, ) time.sleep(delay) continue raise raise RuntimeError("unreachable") # pragma: no cover return wrapper class Database(_Database): @overload def transaction[**P, T](self, func: Callable[P, T]) -> Callable[P, T]: ... @overload def transaction( self, func: None = None, *, commit_on_error: type[Exception] | tuple[type[Exception], ...] | None = None, ) -> _Transaction: ... def transaction[**P, T]( # type: ignore[override] self, func: Callable[P, T] | None = None, *, commit_on_error: type[Exception] | tuple[type[Exception], ...] | None = None, ) -> Callable[P, T] | _Transaction: if func is not None: # Decorator usage: wrap the transaction-bound callable with retry. # Each call opens a fresh transaction, so retrying is safe. return _retry(super().transaction(func)) # Context-manager / no-arg usage: caller owns the retry loop. return super().transaction(commit_on_error=commit_on_error) def get_db() -> Database: url = settings.db_url if settings.dsql_endpoint: 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.""" 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), }