import asyncio import logging from collections import OrderedDict from collections.abc import AsyncIterator from contextlib import asynccontextmanager from dataclasses import dataclass from typing import Dict from uuid import UUID from psycopg import AsyncConnection from psycopg_pool import AsyncConnectionPool from psycopg.rows import dict_row from app.core.config import settings logger = logging.getLogger(__name__) _check_async_connection = AsyncConnectionPool.check_connection @dataclass class PoolEntry: pool: AsyncConnectionPool connection_url: str pool_min_size: int pool_max_size: int borrow_count: int = 0 @dataclass(frozen=True) class CacheKey: project_id: UUID db_role: str class ProjectConnectionManager: def __init__(self) -> None: self._ts_cache: Dict[CacheKey, PoolEntry] = OrderedDict() self._pg_raw_cache: Dict[CacheKey, PoolEntry] = OrderedDict() self._retired_ts: list[tuple[CacheKey, PoolEntry]] = [] self._retired_pg: list[tuple[CacheKey, PoolEntry]] = [] self._ts_lock = asyncio.Lock() self._pg_raw_lock = asyncio.Lock() async def _get_timescale_pool_locked( self, project_id: UUID, db_role: str, connection_url: str, pool_min_size: int, pool_max_size: int, ) -> tuple[CacheKey, AsyncConnectionPool]: pool_min_size = max(0, pool_min_size) pool_max_size = max(1, pool_min_size, pool_max_size) key = CacheKey(project_id=project_id, db_role=db_role) entry = self._ts_cache.get(key) if entry: if ( entry.connection_url == connection_url and entry.pool_min_size == pool_min_size and entry.pool_max_size == pool_max_size ): self._ts_cache.move_to_end(key) return key, entry.pool logger.info( "Rebuilding TimescaleDB pool for project %s (%s) due to config change", project_id, db_role, ) pool = AsyncConnectionPool( conninfo=connection_url, min_size=pool_min_size, max_size=pool_max_size, open=False, kwargs={"row_factory": dict_row}, check=_check_async_connection, ) await pool.open() if entry is not None: if entry.borrow_count: self._retired_ts.append((key, entry)) else: await entry.pool.close() self._ts_cache[key] = PoolEntry( pool=pool, connection_url=connection_url, pool_min_size=pool_min_size, pool_max_size=pool_max_size, ) logger.info("Created TimescaleDB pool for project %s (%s)", project_id, db_role) return key, pool async def _get_pg_pool_locked( self, project_id: UUID, db_role: str, connection_url: str, pool_min_size: int, pool_max_size: int, ) -> tuple[CacheKey, AsyncConnectionPool]: pool_min_size = max(0, pool_min_size) pool_max_size = max(1, pool_min_size, pool_max_size) key = CacheKey(project_id=project_id, db_role=db_role) entry = self._pg_raw_cache.get(key) if entry: if ( entry.connection_url == connection_url and entry.pool_min_size == pool_min_size and entry.pool_max_size == pool_max_size ): self._pg_raw_cache.move_to_end(key) return key, entry.pool logger.info( "Rebuilding PostgreSQL pool for project %s (%s) due to config change", project_id, db_role, ) pool = AsyncConnectionPool( conninfo=connection_url, min_size=pool_min_size, max_size=pool_max_size, open=False, kwargs={"row_factory": dict_row}, check=_check_async_connection, ) await pool.open() if entry is not None: if entry.borrow_count: self._retired_pg.append((key, entry)) else: await entry.pool.close() self._pg_raw_cache[key] = PoolEntry( pool=pool, connection_url=connection_url, pool_min_size=pool_min_size, pool_max_size=pool_max_size, ) logger.info("Created PostgreSQL pool for project %s (%s)", project_id, db_role) return key, pool @asynccontextmanager async def pg_connection( self, project_id: UUID, db_role: str, connection_url: str, pool_min_size: int, pool_max_size: int, ) -> AsyncIterator[AsyncConnection]: async with self._pg_raw_lock: key, pool = await self._get_pg_pool_locked( project_id, db_role, connection_url, pool_min_size, pool_max_size, ) borrowed_entry = self._pg_raw_cache[key] borrowed_entry.borrow_count += 1 await self._evict_pg_raw_if_needed() try: async with pool.connection() as conn: yield conn finally: async with self._pg_raw_lock: borrowed_entry.borrow_count -= 1 if borrowed_entry.borrow_count == 0 and borrowed_entry.pool is not ( self._pg_raw_cache.get(key).pool if key in self._pg_raw_cache else None ): self._retired_pg = [ item for item in self._retired_pg if item[1] is not borrowed_entry ] await borrowed_entry.pool.close() await self._evict_pg_raw_if_needed() @asynccontextmanager async def timescale_connection( self, project_id: UUID, db_role: str, connection_url: str, pool_min_size: int, pool_max_size: int, ) -> AsyncIterator[AsyncConnection]: async with self._ts_lock: key, pool = await self._get_timescale_pool_locked( project_id, db_role, connection_url, pool_min_size, pool_max_size, ) borrowed_entry = self._ts_cache[key] borrowed_entry.borrow_count += 1 await self._evict_ts_if_needed() try: async with pool.connection() as conn: yield conn finally: async with self._ts_lock: borrowed_entry.borrow_count -= 1 if borrowed_entry.borrow_count == 0 and borrowed_entry.pool is not ( self._ts_cache.get(key).pool if key in self._ts_cache else None ): self._retired_ts = [ item for item in self._retired_ts if item[1] is not borrowed_entry ] await borrowed_entry.pool.close() await self._evict_ts_if_needed() async def _evict_ts_if_needed( self, protected_key: CacheKey | None = None ) -> None: while len(self._ts_cache) > settings.PROJECT_TS_CACHE_SIZE: idle = next( ( (key, entry) for key, entry in self._ts_cache.items() if entry.borrow_count == 0 and key != protected_key ), None, ) if idle is None: return key, entry = idle self._ts_cache.pop(key) await entry.pool.close() logger.info( "Evicted TimescaleDB pool for project %s (%s)", key.project_id, key.db_role, ) async def _evict_pg_raw_if_needed( self, protected_key: CacheKey | None = None ) -> None: while len(self._pg_raw_cache) > settings.PROJECT_PG_CACHE_SIZE: idle = next( ( (key, entry) for key, entry in self._pg_raw_cache.items() if entry.borrow_count == 0 and key != protected_key ), None, ) if idle is None: return key, entry = idle self._pg_raw_cache.pop(key) await entry.pool.close() logger.info( "Evicted PostgreSQL pool for project %s (%s)", key.project_id, key.db_role, ) async def close_project( self, project_id: UUID, db_role: str | None = None ) -> bool: """Close this worker's idle pools for a project. Returns ``False`` without interrupting requests when any matching pool is currently borrowed. Callers may retry after those requests complete. """ closed = True for cache, lock, label in ( (self._ts_cache, self._ts_lock, "TimescaleDB"), (self._pg_raw_cache, self._pg_raw_lock, "PostgreSQL"), ): async with lock: keys = [ key for key in cache if key.project_id == project_id and (db_role is None or key.db_role == db_role) ] for key in keys: entry = cache[key] if entry.borrow_count: closed = False continue cache.pop(key) await entry.pool.close() logger.info( "Closed %s pool for project %s (%s)", label, key.project_id, key.db_role, ) for retired, lock in ( (self._retired_ts, self._ts_lock), (self._retired_pg, self._pg_raw_lock), ): async with lock: matches = [ item for item in retired if item[0].project_id == project_id and (db_role is None or item[0].db_role == db_role) ] if any(entry.borrow_count for _key, entry in matches): closed = False for item in matches: key, entry = item if entry.borrow_count: continue retired.remove(item) await entry.pool.close() return closed async def close_all(self) -> None: async with self._ts_lock: for key, entry in list(self._ts_cache.items()): await entry.pool.close() logger.info( "Closed TimescaleDB pool for project %s (%s)", key.project_id, key.db_role, ) self._ts_cache.clear() for _key, entry in self._retired_ts: await entry.pool.close() self._retired_ts.clear() async with self._pg_raw_lock: for key, entry in list(self._pg_raw_cache.items()): await entry.pool.close() logger.info( "Closed PostgreSQL pool for project %s (%s)", key.project_id, key.db_role, ) self._pg_raw_cache.clear() for _key, entry in self._retired_pg: await entry.pool.close() self._retired_pg.clear() project_connection_manager = ProjectConnectionManager()