Files
TJWaterServerBinary/app/infra/db/dynamic_manager.py
T

347 lines
12 KiB
Python

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()