import logging from asyncio import current_task from time import sleep from sqlalchemy import create_engine from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import create_async_engine, async_scoped_session from sqlalchemy.orm import sessionmaker from sqlalchemy.pool import QueuePool, AsyncAdaptedQueuePool from sqlalchemy.sql import select from config import ( POSTGRES_RECONNECTION_TIME, TRY_TO_RECONNECT, ) from server.db.db_connector import DBConnector log = logging.getLogger(__name__) class PostgresConnector(DBConnector): def __init__(self, user, password, host, port, database): super().__init__(user, password, host, port, database) self._create_sync_session() def _get_async_connect_url(self): return f"postgresql+asyncpg://{self.user}:{self.password}@{self.host}:{self.port}/{self.database}" def _get_sync_connect_url(self): return f"postgresql+psycopg2://{self.user}:{self.password}@{self.host}:{self.port}/{self.database}" def _create_engine(self, flag="sync"): create_engine_mapper = {"sync": create_engine, "async": create_async_engine} return create_engine_mapper[flag]( self._get_sync_connect_url() if flag == "sync" else self._get_async_connect_url(), poolclass=QueuePool if flag == "sync" else AsyncAdaptedQueuePool, pool_size=30, ) def _create_sync_session(self): try: i = 0 engine = self._create_engine() while i != TRY_TO_RECONNECT: try: self.check_db_connection(engine) except Exception as e: log.error(f"Exception caught - {e}") log.warning(f"Reconnecting in {POSTGRES_RECONNECTION_TIME} seconds") sleep(POSTGRES_RECONNECTION_TIME) i += 1 continue else: break else: raise Exception("Retries limit was reached") except Exception as e: log.error(f"Can't connect to db due to {e}") raise e else: log.info("Connected to db via sync engine") def create_async_session(self): try: engine = self._create_engine(flag="async") except Exception as e: log.error(f"Can't connect to db due to {e}") else: log.info("Connected to db via async engine") return async_scoped_session( sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession), scopefunc=current_task ) @staticmethod def check_db_connection(engine): with engine.begin() as conn: conn.execute(select(1)) def get_postgres_async_session(user: str, password: str, host: str, database: str, port: int = 5432): return PostgresConnector( user=user, password=password, host=host, port=port, database=database, ).create_async_session()